//! 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 { 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 { 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) -> Result { 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::(s).unwrap()) .collect::>(); 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::>(); 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::>(); 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| -> Result { 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::>(); 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::(); 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(()) }