Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,499 @@
|
||||
//! MTPLX runtime vector SDPA, BF16/D256 on the evaluation M5 Max (applegpu_g17s).
|
||||
//! Original 1/2-pass kernels. Full attention and surrounding model routing are open.
|
||||
use super::super::gpu::{Buffer, MetalConstant};
|
||||
use super::array::{Array, Dtype, Layout};
|
||||
use super::stream::{Device, Stream};
|
||||
use super::{bytes, dispatch_specialized, encoder, ops, tensor, views};
|
||||
|
||||
pub(super) fn vector(
|
||||
q: &Array,
|
||||
k: &Array,
|
||||
v: &Array,
|
||||
mask: Option<&Array>,
|
||||
scale: f32,
|
||||
causal: bool,
|
||||
stream: Stream,
|
||||
) -> Result<Array, String> {
|
||||
let qs = q.layout().shape().to_vec();
|
||||
let ks = k.layout().shape().to_vec();
|
||||
if stream.device() != Device::Gpu
|
||||
|| qs.len() != 4
|
||||
|| ks.len() != 4
|
||||
|| k.layout().shape() != v.layout().shape()
|
||||
|| qs[0] != ks[0]
|
||||
|| qs[3] != 256
|
||||
|| ks[3] != 256
|
||||
|| ks[1] <= 0
|
||||
|| qs[1] % ks[1] != 0
|
||||
|| !(1..=8).contains(&qs[2])
|
||||
|| qs[2] > ks[2]
|
||||
|| qs[2] * (qs[1] / ks[1]) > 32
|
||||
|| [q, k, v]
|
||||
.iter()
|
||||
.any(|a| a.layout().dtype() != Dtype::BF16 || a.layout().size() == 0)
|
||||
|| !scale.is_finite()
|
||||
|| (causal && mask.is_some())
|
||||
{
|
||||
return Err(
|
||||
"vector SDPA requires compatible BF16/D256 GPU inputs, S<=8 and S*GQA<=32".into(),
|
||||
);
|
||||
}
|
||||
let mut inputs = vec![q.clone(), k.clone(), v.clone()];
|
||||
if let Some(mask) = mask {
|
||||
let dtype = mask.layout().dtype();
|
||||
if dtype.promote(Dtype::BF16) != Dtype::BF16 {
|
||||
return Err("SDPA mask must promote to BF16".into());
|
||||
}
|
||||
let mask = ops::astype(
|
||||
mask,
|
||||
if dtype == Dtype::Bool {
|
||||
dtype
|
||||
} else {
|
||||
Dtype::BF16
|
||||
},
|
||||
false,
|
||||
stream,
|
||||
)?;
|
||||
inputs.push(ops::broadcast_to(
|
||||
&mask,
|
||||
&[qs[0], qs[1], qs[2], ks[2]],
|
||||
stream,
|
||||
)?);
|
||||
}
|
||||
Array::make_operation_with_inputs(
|
||||
stream,
|
||||
&inputs,
|
||||
Layout::new(&qs, Dtype::BF16)?,
|
||||
ops::Operation::SdpaVector { scale, causal },
|
||||
)
|
||||
}
|
||||
|
||||
fn copy(input: &Array) -> Result<Array, String> {
|
||||
let layout = input.layout();
|
||||
let out = Array::new(
|
||||
layout.shape(),
|
||||
layout.dtype(),
|
||||
Buffer::mtplx_bytes(layout.nbytes() as u64)?,
|
||||
)?;
|
||||
drop(layout);
|
||||
views::general_copy_inplace(input, &out)?;
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn blocks(n: i32, simds: i32, override_value: i32, arch: Option<char>) -> Result<i32, String> {
|
||||
let default = if arch == Some('d') {
|
||||
if simds <= 2 && n > 8192 {
|
||||
256
|
||||
} else if simds >= 6 && n >= 65536 {
|
||||
1024
|
||||
} else if simds >= 6 && n >= 16384 {
|
||||
512
|
||||
} else {
|
||||
128
|
||||
}
|
||||
} else if arch != Some('s') {
|
||||
if simds >= 4 { 64 } else { 32 }
|
||||
} else if n <= 1024 || simds <= 4 {
|
||||
64
|
||||
} else if n <= 8192 {
|
||||
128
|
||||
} else if n <= 32768 {
|
||||
256
|
||||
} else if n <= 65536 {
|
||||
512
|
||||
} else {
|
||||
1024
|
||||
};
|
||||
if override_value > 0 {
|
||||
override_value
|
||||
.checked_add(31)
|
||||
.map(|v| v / 32 * 32)
|
||||
.ok_or("SDPA block count overflow".into())
|
||||
} else {
|
||||
Ok(default)
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires Apple M5 Max; actual pinned fused SDPA kernels"]
|
||||
fn mtplx_sdpa_vector_matches_reference() {
|
||||
use super::super::{
|
||||
gpu::Context,
|
||||
qwen_mtplx_tests::{capture_dispatches, pattern, verify},
|
||||
};
|
||||
use super::{allocator, configure_sources, stream::Streams};
|
||||
configure_sources().unwrap();
|
||||
let _context = Context::open_qwen(0).unwrap();
|
||||
let streams = Streams::new(Some(allocator::Allocator::new().unwrap()));
|
||||
let gpu = streams.default_stream(Device::Gpu).unwrap();
|
||||
assert!(
|
||||
std::env::var("MLX_SDPA_BLOCKS").is_err(),
|
||||
"reference receipt uses default blocks"
|
||||
);
|
||||
for (n, s, over, want) in [
|
||||
(1024, 24, 0, 64),
|
||||
(16383, 24, 0, 256),
|
||||
(16384, 24, 0, 256),
|
||||
(65535, 24, 0, 512),
|
||||
(65536, 24, 0, 512),
|
||||
(65537, 24, 0, 1024),
|
||||
(8193, 2, 0, 64),
|
||||
(1, 1, 33, 64),
|
||||
(1, 1, 31, 32),
|
||||
] {
|
||||
assert_eq!(blocks(n, s, over, Some('s')).unwrap(), want);
|
||||
}
|
||||
let cases = include_str!("../../../../tests/fixtures/mtplx-sdpa-vector.jsonl")
|
||||
.lines()
|
||||
.map(|s| serde_json::from_str::<serde_json::Value>(s).unwrap())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(cases.len(), 48);
|
||||
let bf16 = |shape: &[i32], salt| {
|
||||
Array::new(
|
||||
shape,
|
||||
Dtype::BF16,
|
||||
pattern(shape.iter().map(|&v| v as u32).product(), salt),
|
||||
)
|
||||
.unwrap()
|
||||
};
|
||||
for prepared in [false, true] {
|
||||
for case in &cases {
|
||||
super::eval::execution_tests::observed(|| {
|
||||
let [batch, heads, rows, total] =
|
||||
["batch", "heads", "rows", "total"].map(|k| case[k].as_i64().unwrap() as i32);
|
||||
let step = if case["copied"].as_bool().unwrap() {
|
||||
2
|
||||
} else {
|
||||
1
|
||||
};
|
||||
let mode = case["mode"].as_str().unwrap();
|
||||
let q = ops::transpose(
|
||||
&ops::slice(
|
||||
&bf16(&[batch, rows, heads, 256 * step], 103),
|
||||
&[0, 0, 0, 0],
|
||||
&[batch, rows, heads, 256 * step],
|
||||
&[1, 1, 1, step],
|
||||
gpu,
|
||||
)
|
||||
.unwrap(),
|
||||
&[0, 2, 1, 3],
|
||||
gpu,
|
||||
)
|
||||
.unwrap();
|
||||
let kv = |salt| {
|
||||
ops::slice(
|
||||
&bf16(&[batch, 2, total + 7, 256 * step], salt),
|
||||
&[0, 0, 2, 0],
|
||||
&[batch, 2, total + 2, 256 * step],
|
||||
&[1, 1, 1, step],
|
||||
gpu,
|
||||
)
|
||||
.unwrap()
|
||||
};
|
||||
let (k, v) = (kv(104), kv(105));
|
||||
let mask = if matches!(mode, "bool" | "add") {
|
||||
let data = (0..rows * total)
|
||||
.flat_map(|i| {
|
||||
if mode == "bool" {
|
||||
vec![u8::from(i % 3 != 1)]
|
||||
} else {
|
||||
(if i % 3 != 1 { 0u16 } else { 0xc120 })
|
||||
.to_le_bytes()
|
||||
.to_vec()
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Some(
|
||||
Array::new(
|
||||
&[1, 1, rows, total],
|
||||
if mode == "bool" {
|
||||
Dtype::Bool
|
||||
} else {
|
||||
Dtype::BF16
|
||||
},
|
||||
super::scalar_buffer(&data).unwrap(),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if prepared {
|
||||
ops::evaluate(&streams, &[q.clone(), k.clone(), v.clone()], gpu, false)
|
||||
.unwrap();
|
||||
}
|
||||
let out = if mode == "causal" {
|
||||
vector(&q, &k, &v, None, 0.0625, true, gpu)
|
||||
} else {
|
||||
super::attention::scaled_dot_product(&q, &k, &v, mask.as_ref(), 0.0625, gpu)
|
||||
}
|
||||
.unwrap();
|
||||
drop((q, k, v, mask));
|
||||
let (_, calls) = capture_dispatches(|| {
|
||||
ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap()
|
||||
});
|
||||
let names = calls
|
||||
.iter()
|
||||
.filter(|r| r.0.starts_with("sdpa_vector"))
|
||||
.map(|r| r.0.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
names,
|
||||
if total >= 1024 {
|
||||
vec![
|
||||
"sdpa_vector_2pass_1_bfloat16_t_256_256",
|
||||
"sdpa_vector_2pass_2_bfloat16_t_256",
|
||||
]
|
||||
} else {
|
||||
vec!["sdpa_vector_bfloat16_t_256_256"]
|
||||
}
|
||||
);
|
||||
let mut data = vec![0; out.layout().nbytes()];
|
||||
out.buffer().read(out.offset(), &mut data).unwrap();
|
||||
verify(
|
||||
case["kernel"].as_str().unwrap(),
|
||||
&[(
|
||||
&super::scalar_buffer(&data).unwrap(),
|
||||
out.layout().size() as u32,
|
||||
)],
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
streams.clear_streams().unwrap();
|
||||
}
|
||||
|
||||
pub(super) fn evaluate(
|
||||
inputs: &[Array],
|
||||
out: &Array,
|
||||
scale: f32,
|
||||
causal: bool,
|
||||
) -> Result<(), String> {
|
||||
let qp = &inputs[0];
|
||||
let ql = qp.layout();
|
||||
let shape = ql.shape().to_vec();
|
||||
let st = ql.strides();
|
||||
let bidx = if shape[0] == 1 { 1 } else { 0 };
|
||||
let q_ok = ql.flags().row_contiguous
|
||||
|| ((shape[0] == 1 || shape[1] == 1)
|
||||
&& st[3] == 1
|
||||
&& st[2] == i64::from(shape[3]) * i64::from(shape[bidx])
|
||||
&& st[bidx] == i64::from(shape[3]));
|
||||
drop(ql);
|
||||
let q = if q_ok { qp.clone() } else { copy(qp)? };
|
||||
let mut copies = Vec::new();
|
||||
let copy_unless =
|
||||
|a: &Array, predicate: bool, copies: &mut Vec<Array>| -> Result<Array, String> {
|
||||
if predicate {
|
||||
Ok(a.clone())
|
||||
} else {
|
||||
let c = copy(a)?;
|
||||
copies.push(c.clone());
|
||||
Ok(c)
|
||||
}
|
||||
};
|
||||
let kv_ok = |a: &Array| {
|
||||
let l = a.layout();
|
||||
let s = l.shape();
|
||||
let t = l.strides();
|
||||
t[3] == 1 && (s[0] == 1 || s[1] == 1 || t[0] == t[1] * i64::from(s[1]))
|
||||
};
|
||||
let k = copy_unless(&inputs[1], kv_ok(&inputs[1]), &mut copies)?;
|
||||
let v = copy_unless(&inputs[2], kv_ok(&inputs[2]), &mut copies)?;
|
||||
// A non-copied q has the extra descriptor held by the primitive, just as
|
||||
// the original's local `array q` does. Only a unique copied query donates.
|
||||
let ql = q.layout();
|
||||
if q.is_donatable() && ql.flags().row_contiguous && ql.size() == out.layout().size() {
|
||||
out.copy_shared_buffer(&q, ql.strides(), ql.flags(), q.data_size(), 0)?;
|
||||
} else {
|
||||
if !q_ok {
|
||||
copies.push(q.clone());
|
||||
}
|
||||
let nbytes = out.layout().nbytes();
|
||||
out.set_data(Buffer::mtplx_bytes(nbytes as u64)?)?;
|
||||
}
|
||||
drop(ql);
|
||||
let mask = if let Some(m) = inputs.get(3) {
|
||||
let l = m.layout();
|
||||
let s = l.shape();
|
||||
let t = l.strides();
|
||||
let ok = l.flags().row_contiguous
|
||||
|| shape[0] == 1
|
||||
|| shape[1] == 1
|
||||
|| t[0] == t[1] * i64::from(s[1]);
|
||||
drop(l);
|
||||
Some(copy_unless(m, ok, &mut copies)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let kl = k.layout();
|
||||
let vl = v.layout();
|
||||
let n = kl.shape()[2];
|
||||
let gqa = shape[1] / kl.shape()[1];
|
||||
let kh = kl.strides()[if kl.shape()[1] == 1 { 0 } else { 1 }] as u64;
|
||||
let ks = kl.strides()[2] as u64;
|
||||
let vh = vl.strides()[if vl.shape()[1] == 1 { 0 } else { 1 }] as u64;
|
||||
let vs = vl.strides()[2] as u64;
|
||||
let arch = encoder::arch_suffix();
|
||||
let two = (matches!(arch, Some('s' | 'd')) && n >= 1024) || (gqa > 1 && n >= 4096);
|
||||
let bool_mask = mask
|
||||
.as_ref()
|
||||
.is_some_and(|a| a.layout().dtype() == Dtype::Bool);
|
||||
let mut constants = [
|
||||
(20, mask.is_some()),
|
||||
(21, !q.layout().flags().row_contiguous),
|
||||
(22, causal && shape[2] > 1),
|
||||
(23, bool_mask),
|
||||
(24, mask.is_some() && !bool_mask),
|
||||
(25, false),
|
||||
]
|
||||
.map(|(index, value)| MetalConstant {
|
||||
index,
|
||||
value: u32::from(value),
|
||||
kind: 0,
|
||||
})
|
||||
.into_iter()
|
||||
.collect::<Vec<_>>();
|
||||
let mask_layout = mask.as_ref().map(Array::layout);
|
||||
let mask_strides = mask_layout
|
||||
.as_ref()
|
||||
.map(|l| -> Result<[i32; 3], String> {
|
||||
let s = l.shape();
|
||||
let t = l.strides();
|
||||
let raw = [
|
||||
if s[3] > 1 { t[3] } else { 0 },
|
||||
if s[2] > 1 { t[2] } else { 0 },
|
||||
if s[1] > 1 {
|
||||
t[1]
|
||||
} else if s[0] > 1 {
|
||||
t[0]
|
||||
} else {
|
||||
0
|
||||
},
|
||||
];
|
||||
let [a, b, c] = raw.map(|v| i32::try_from(v).map_err(|_| "SDPA mask stride overflow"));
|
||||
Ok([a?, b?, c?])
|
||||
})
|
||||
.transpose()?;
|
||||
let qb = q.buffer();
|
||||
let kb = k.buffer();
|
||||
let vb = v.buffer();
|
||||
let ob = out.buffer();
|
||||
let mb = mask.as_ref().map(Array::buffer);
|
||||
let base = vec![
|
||||
tensor(0, &qb)
|
||||
.at_byte_offset(q.offset())
|
||||
.input(q.data_size()),
|
||||
tensor(1, &kb)
|
||||
.at_byte_offset(k.offset())
|
||||
.input(k.data_size()),
|
||||
tensor(2, &vb)
|
||||
.at_byte_offset(v.offset())
|
||||
.input(v.data_size()),
|
||||
];
|
||||
let mask_bindings = |offset: u32| {
|
||||
let mut bindings = Vec::new();
|
||||
if let (Some(m), Some(ms)) = (&mb, &mask_strides) {
|
||||
bindings.push(
|
||||
tensor(offset + u32::from(!bool_mask), m)
|
||||
.at_byte_offset(mask.as_ref().unwrap().offset())
|
||||
.input(mask.as_ref().unwrap().data_size()),
|
||||
);
|
||||
bindings.extend([
|
||||
bytes(offset + 2, &ms[0]),
|
||||
bytes(offset + 3, &ms[1]),
|
||||
bytes(offset + 4, &ms[2]),
|
||||
]);
|
||||
}
|
||||
bindings
|
||||
};
|
||||
if two {
|
||||
let count = blocks(
|
||||
n,
|
||||
gqa * shape[2],
|
||||
{
|
||||
let value = unsafe { libc::getenv(c"MLX_SDPA_BLOCKS".as_ptr()) };
|
||||
if value.is_null() {
|
||||
0
|
||||
} else {
|
||||
unsafe { libc::atoi(value) }
|
||||
}
|
||||
},
|
||||
arch,
|
||||
)?;
|
||||
constants.push(MetalConstant {
|
||||
index: 26,
|
||||
value: count as u32,
|
||||
kind: 1,
|
||||
});
|
||||
let rows = shape[..3].iter().map(|&v| v as u64).product::<u64>();
|
||||
let temporaries = [
|
||||
Buffer::mtplx_bytes(rows * count as u64 * 256 * 2)?,
|
||||
Buffer::mtplx_bytes(rows * count as u64 * 4)?,
|
||||
Buffer::mtplx_bytes(rows * count as u64 * 4)?,
|
||||
];
|
||||
encoder::add_temporaries(&temporaries);
|
||||
let [partial, sums, maxs] = &temporaries;
|
||||
let mut bindings = base;
|
||||
bindings.extend([
|
||||
tensor(3, partial).output((rows * count as u64 * 256) as usize),
|
||||
tensor(4, sums).output((rows * count as u64) as usize),
|
||||
tensor(5, maxs).output((rows * count as u64) as usize),
|
||||
bytes(7, &n),
|
||||
bytes(8, &kh),
|
||||
bytes(9, &ks),
|
||||
bytes(10, &vh),
|
||||
bytes(11, &vs),
|
||||
bytes(12, &scale),
|
||||
]);
|
||||
bindings.extend(mask_bindings(13));
|
||||
dispatch_specialized(
|
||||
"sdpa_vector_2pass_1_bfloat16_t_256_256",
|
||||
&bindings,
|
||||
&constants,
|
||||
[kl.shape()[1] as u32, shape[0] as u32, count as u32],
|
||||
[32, gqa as u32, shape[2] as u32],
|
||||
)?;
|
||||
dispatch_specialized(
|
||||
"sdpa_vector_2pass_2_bfloat16_t_256",
|
||||
&[
|
||||
tensor(0, partial).input((rows * count as u64 * 256) as usize),
|
||||
tensor(1, sums).input((rows * count as u64) as usize),
|
||||
tensor(2, maxs).input((rows * count as u64) as usize),
|
||||
tensor(3, &ob)
|
||||
.at_byte_offset(out.offset())
|
||||
.output(out.data_size()),
|
||||
bytes(4, &count),
|
||||
],
|
||||
&[],
|
||||
[(shape[0] * shape[1]) as u32, shape[2] as u32, 1],
|
||||
[1024, 1, 1],
|
||||
)?;
|
||||
} else {
|
||||
let mut bindings = base;
|
||||
bindings.extend([
|
||||
tensor(3, &ob)
|
||||
.at_byte_offset(out.offset())
|
||||
.output(out.data_size()),
|
||||
bytes(4, &gqa),
|
||||
bytes(5, &n),
|
||||
bytes(6, &kh),
|
||||
bytes(7, &ks),
|
||||
bytes(8, &vh),
|
||||
bytes(9, &vs),
|
||||
bytes(10, &scale),
|
||||
]);
|
||||
bindings.extend(mask_bindings(11));
|
||||
dispatch_specialized(
|
||||
"sdpa_vector_bfloat16_t_256_256",
|
||||
&bindings,
|
||||
&constants,
|
||||
[(shape[0] * shape[1]) as u32, shape[2] as u32, 1],
|
||||
[1024, 1, 1],
|
||||
)?;
|
||||
}
|
||||
for copy in &copies {
|
||||
encoder::add_temporaries(std::slice::from_ref(&*copy.buffer()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
Reference in New Issue
Block a user