Files
DS4Server/src/engine/metal/qwen_mtplx/sdpa.rs
T

500 lines
16 KiB
Rust

//! 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(())
}