//! Floating Matmul on the evaluation M5 Max: projections and attention batches. //! Original broadcasting, batch collapse, transpose preparation and dispatch. //! Original vector/empty/dtype rules; not complex/other-device matrix arithmetic. use super::super::gpu::{Buffer, MetalConstant}; use super::array::{Array, Dtype, Layout}; use super::stream::{Device, Stream}; use super::{bytes, dispatch_geometry, dispatch_specialized, encoder, ops, tensor, views}; pub(super) fn matmul(a: &Array, b: &Array, stream: Stream) -> Result { let ashape = a.layout().shape().to_vec(); let bshape = b.layout().shape().to_vec(); let dtype = a.layout().dtype().promote(b.layout().dtype()); if stream.device() != Device::Gpu || ashape.is_empty() || bshape.is_empty() || !matches!(dtype, Dtype::BF16 | Dtype::F16 | Dtype::F32) { return Err("matmul requires nonscalar GPU inputs promoting to BF16/F16/F32".into()); } let a = if ashape.len() == 1 { ops::expand_dims(a, &[0], stream)? } else { a.clone() }; let b = if bshape.len() == 1 { ops::expand_dims(b, &[1], stream)? } else { b.clone() }; if a.layout().dim(-1)? != b.layout().dim(-2)? { return Err("matmul inner dimensions differ".into()); } let n = b.layout().dim(-1)?; let a = ops::astype(&a, dtype, false, stream)?; let b = ops::astype(&b, dtype, false, stream)?; let flattened = ashape.len() > 2 && bshape.len() <= 2; let inputs = if flattened { let rows = ashape[..ashape.len() - 1] .iter() .try_fold(1i32, |n, &d| n.checked_mul(d)) .ok_or("matmul flattened dimension overflow")?; vec![ Array::make_operation( stream, &a, Layout::new(&[rows, *ashape.last().unwrap()], dtype)?, ops::Operation::Flatten { start: 0, end: ashape.len() - 2, }, )?, b, ] } else if ashape.len() > 2 || bshape.len() > 2 { let rank = a.layout().shape().len().max(b.layout().shape().len()); ops::broadcast_arrays(&[a, b], &[rank - 2, rank - 1], stream)? } else { vec![a, b] }; let mut shape = inputs[0].layout().shape().to_vec(); *shape.last_mut().unwrap() = n; let mut out = Array::make_operation_with_inputs( stream, &inputs, Layout::new(&shape, dtype)?, ops::Operation::DenseMatmul, )?; if flattened { let mut shape = ashape.clone(); *shape.last_mut().unwrap() = n; out = Array::make_operation( stream, &out, Layout::new(&shape, dtype)?, ops::Operation::Unflatten { axis: 0, shape: ashape[..ashape.len() - 1].to_vec(), }, )?; } let mut axes = Vec::new(); if ashape.len() == 1 { axes.push(out.layout().shape().len() as i32 - 2); } if bshape.len() == 1 { axes.push(out.layout().shape().len() as i32 - 1); } if axes.is_empty() { Ok(out) } else { ops::squeeze(&out, Some(&axes), stream) } } pub(super) fn linear(input: &Array, weight: &Array, stream: Stream) -> Result { // nn.Linear delegates to the same Matmul factory (including vectors, // promotion, empty matrices and leading-axis flatten/unflatten). matmul(input, &ops::transpose(weight, &[1, 0], stream)?, stream) } fn prepare( input: &Array, vector: bool, copies: &mut Vec, ) -> Result<(bool, i32, Array), String> { let layout = input.layout(); let rank = layout.shape().len(); let [rows, cols]: [i32; 2] = layout.shape()[rank - 2..] .try_into() .map_err(|_| "dense matmul rank")?; let [sx, sy]: [i64; 2] = layout.strides()[rank - 2..] .try_into() .map_err(|_| "dense matmul strides")?; if sy == 1 && (!vector || sx == i64::from(cols)) { return Ok(( false, i32::try_from(sx).map_err(|_| "dense leading dimension overflow")?, input.clone(), )); } if sx == 1 && (!vector || sy == i64::from(rows)) { return Ok(( true, i32::try_from(sy).map_err(|_| "dense leading dimension overflow")?, input.clone(), )); } let copy = Array::new( layout.shape(), layout.dtype(), Buffer::mtplx_bytes(layout.nbytes() as u64)?, )?; drop(layout); views::general_copy_inplace(input, ©)?; copies.push(copy.clone()); Ok((false, cols, copy)) } pub(super) fn tf32_enabled() -> bool { static ENABLED: std::sync::OnceLock = std::sync::OnceLock::new(); *ENABLED.get_or_init(|| { let value = unsafe { libc::getenv(c"MLX_ENABLE_TF32".as_ptr()) }; value.is_null() || unsafe { libc::atoi(value) != 0 } }) } fn dot(a: &Array, b: &Array, out: &Array, k: u32) -> Result<(), String> { let blocks = k.div_ceil(32 * 512); let partials = Array::new( &[blocks as i32], Dtype::F32, Buffer::mtplx_bytes(u64::from(blocks) * 4)?, )?; dispatch_geometry( &format!( "dot_product_{}_it32_tg512_sg16", a.layout().dtype().kernel_name() ), &[ tensor(0, &a.buffer()) .at_byte_offset(a.offset()) .input(a.data_size()), tensor(1, &b.buffer()) .at_byte_offset(b.offset()) .input(b.data_size()), tensor(2, &partials.buffer()).output(blocks as usize), bytes(3, &k), ], &[], [blocks * 512, 1, 1], [512, 1, 1], true, )?; let temp = Array::new(out.layout().shape(), Dtype::F32, Buffer::mtplx_bytes(4)?)?; super::reduce::all(&partials, &temp)?; views::copy_gpu(&temp, out, false)?; for a in [&partials, &temp] { encoder::add_temporaries(std::slice::from_ref(&*a.buffer())); } Ok(()) } pub(super) fn evaluate(a: &Array, b: &Array, out: &Array) -> Result<(), String> { let dtype = out.layout().dtype(); let dtype_name = dtype.kernel_name(); if a.layout().size() == 0 || b.layout().size() == 0 { let zero = Array::new( &[], dtype, super::scalar_buffer(&vec![0; dtype.itemsize()])?, )?; views::full(&zero, out)?; encoder::add_temporaries(std::slice::from_ref(&*zero.buffer())); return Ok(()); } let mut m = a.layout().dim(-2)? as u32; let n = b.layout().dim(-1)? as u32; let k = a.layout().dim(-1)? as u32; let nbytes = out.layout().nbytes(); out.set_data(Buffer::mtplx_bytes(nbytes as u64)?)?; let mut copies = Vec::new(); let (ta, lda, a) = prepare(a, m == 1, &mut copies)?; let (tb, ldb, b) = prepare(b, n == 1, &mut copies)?; let (mut batch_shape, mut strides) = { let al = a.layout(); let bl = b.layout(); let rank = al.shape().len(); views::collapse_joint( &al.shape()[..rank - 2], &[&al.strides()[..rank - 2], &bl.strides()[..rank - 2]], ) }; if batch_shape.is_empty() { batch_shape.push(1); strides = vec![vec![0], vec![0]]; } let mut batch = u32::try_from(out.layout().size() / (m as usize * n as usize)) .map_err(|_| "matmul batch overflow")?; if batch > 1 && !ta && batch_shape.len() == 1 && a.layout().stride(-2)? == i64::from(k) && strides[0][0] == i64::from(m) * i64::from(k) && strides[1][0] == 0 { m = m .checked_mul(batch_shape[0] as u32) .filter(|&v| v <= i32::MAX as u32) .ok_or("flattened matmul M overflow")?; batch = 1; batch_shape = vec![1]; strides = vec![vec![0], vec![0]]; } let [a_batch, b_batch]: [Vec; 2] = strides.try_into().unwrap(); let batch_ndim = batch_shape.len() as i32; if m == 1 && n == 1 && batch == 1 && a.layout().flags().row_contiguous && b.layout().flags().row_contiguous { dot(&a, &b, out, k)?; for copy in copies { encoder::add_temporaries(std::slice::from_ref(&*copy.buffer())); } return Ok(()); } let ab = a.buffer(); let bb = b.buffer(); let ob = out.buffer(); let ai = |index| { tensor(index, &ab) .at_byte_offset(a.offset()) .input(a.data_size()) }; let bi = |index| { tensor(index, &bb) .at_byte_offset(b.offset()) .input(b.data_size()) }; let output = |index| tensor(index, &ob).output(out.data_size()); let constants = |pairs: &[(u32, bool)]| { pairs .iter() .map(|&(index, value)| MetalConstant { index, value: u32::from(value), kind: 0, }) .collect::>() }; let passes = m.div_ceil(5); if !ta && tb && dtype != Dtype::F32 && n > 1 && (2..=15).contains(&m) && k.is_multiple_of(4) && a.offset().is_multiple_of(8) && b.offset().is_multiple_of(8) && lda % 4 == 0 && ldb % 4 == 0 && a_batch.iter().chain(&b_batch).all(|s| s % 4 == 0) { let vectors = m.div_ceil(passes); let lanes = if passes == 1 || n <= 64 { 32 } else { 16 }; dispatch_specialized( &format!("gemv_wide_{dtype_name}_nv{vectors}_kl{lanes}"), &[ bi(0), ai(1), output(3), bytes(4, &k), bytes(5, &n), bytes(6, &m), bytes(7, &ldb), bytes(8, &lda), bytes(11, &batch_ndim), bytes(12, batch_shape.as_slice()), bytes(13, a_batch.as_slice()), bytes(14, b_batch.as_slice()), ], &constants(&[(0, batch_ndim != 1), (1, false)]), [if n >= 65536 { 1 } else { passes }, n.div_ceil(4), batch], [32, lanes / 8, 1], )?; } else if m.min(n) == 1 { let b_matrix = n != 1; let transposed = if b_matrix { !tb } else { ta }; let length = if b_matrix { n } else { m }; let ld = if b_matrix { ldb } else { lda }; let (mat_batch, vec_batch) = if b_matrix { (&b_batch, &a_batch) } else { (&a_batch, &b_batch) }; let (mut bm, mut bn, mut sm, mut sn, mut tm, mut tn) = (1, 1, 1, 32, 4, 4); let per_group; if transposed { (sm, sn) = if k >= 8192 && length >= 2048 { (4, 8) } else { (8, 4) }; bn = if length >= 2048 { 16 } else if length >= 512 { 4 } else { 2 }; if length < tn { tn = 1; } per_group = bn * sn * tn; } else { bm = if length >= 4096 { 8 } else { 4 }; if k <= 64 { bm = 1; sm = 8; sn = 4; } else if u64::from(k) >= 16 * u64::from(length) { bm = 1; bn = 8; } if length < tm { tm = 1; } per_group = bm * sm * tm; } dispatch_specialized( &format!( "gemv_{}{dtype_name}_bm{bm}_bn{bn}_sm{sm}_sn{sn}_tm{tm}_tn{tn}_nc{}_axpby0", if transposed { "t_" } else { "" }, u8::from(batch_ndim != 1) ), &[ if b_matrix { bi(0) } else { ai(0) }, if b_matrix { ai(1) } else { bi(1) }, output(3), bytes(4, &k), bytes(5, &length), bytes(6, &ld), bytes(9, &batch_ndim), bytes(10, batch_shape.as_slice()), bytes(11, vec_batch.as_slice()), bytes(12, mat_batch.as_slice()), ], &[], [length.div_ceil(per_group), 1, batch], [32, bn, bm], )?; } else { let trans = format!( "{}{}", if ta { 't' } else { 'n' }, if tb { 't' } else { 'n' } ); // M5 Max (architecture suffix s), matching the runtime's cached flag. let nax = tf32_enabled() || dtype != Dtype::F32; let split = batch == 1 && if nax { u64::from(k) >= 3 * u64::from(m.max(n)) || (m.max(n) <= 1024 && k > 2 * m.max(n)) } else { u64::from(m.div_ceil(16)) * u64::from(n.div_ceil(16)) <= 2048 && k / 16 >= 8 && k >= m.max(n) }; if split { let (bm, bn, bk, wm, wn) = if !nax { ( if m < 40 { 16 } else { 32 }, if n < 40 { 16 } else { 32 }, 16, 2, 2, ) } else if (u64::from(m) + u64::from(n)) / 2 < 512 || k <= 4096 { (64, 64, 256, 2, 2) } else { (128, 128, 512, 4, 4) }; let simd_parts = (k / 16 / (m.div_ceil(32) * n.div_ceil(32))) .next_power_of_two() .clamp(2, 32); let part_size = if !nax { (k / 16 / simd_parts) * 16 } else if k <= 1024 { k / 2 } else if k <= 2048 { 1024 } else if k <= 4096 { 2048 } else { 4096 }; let parts = if nax { k.div_ceil(part_size) } else { simd_parts }; let stride = m .checked_mul(n) .filter(|&v| v <= i32::MAX as u32) .ok_or("dense split output overflow")?; let intermediate = Buffer::mtplx_bytes(u64::from(parts) * u64::from(stride) * 4)?; let (tn, tm) = (n.div_ceil(bn), m.div_ceil(bm)); let swizzle = u32::from(nax && tm > 3); let params = [ m, n, k, lda as u32, ldb as u32, n, tn, tm, parts, stride, part_size, swizzle, part_size / bk, ]; let name = if nax { format!( "steel_gemm_splitk_nax_{trans}_{dtype_name}_float32_bm{bm}_bn{bn}_bk{bk}_wm{wm}_wn{wn}" ) } else { format!( "steel_gemm_splitk_{trans}_{dtype_name}_float32_bm{bm}_bn{bn}_bk{bk}_wm{wm}_wn{wn}_MN_{}aligned_K_{}aligned", if m.is_multiple_of(bm) && n.is_multiple_of(bn) { 't' } else { 'n' }, if k.is_multiple_of(bk) { 't' } else { 'n' } ) }; dispatch_specialized( &name, &[ ai(0), bi(1), tensor(2, &intermediate).output(parts as usize * stride as usize), bytes(3, ¶ms), ], &if nax { constants(&[(200, m.is_multiple_of(bm)), (201, n.is_multiple_of(bn))]) } else { vec![] }, if nax { [ tn * (1 << swizzle) * tm.div_ceil(1 << swizzle) * parts, 1, 1, ] } else { [tn, tm, parts] }, [32, wn, wm], )?; dispatch_geometry( &format!("steel_gemm_splitk_accum_{dtype_name}_float32"), &[ tensor(0, &intermediate).input(parts as usize * stride as usize), output(1), bytes(2, &parts), bytes(3, &stride), bytes(4, &n), ], &[], [n, m, 1], super::block_dims([n, m, 1]), true, )?; encoder::add_temporaries(std::slice::from_ref(&intermediate)); } else { let bk = if !nax { 16 } else if k >= 8192 && u64::from(k) > u64::from(m) + u64::from(n) { 64 } else { 256 }; let (bn, wn, swizzle) = if nax { (128, 4, 2) } else { (64, 2, 0) }; let (tn, tm) = (n.div_ceil(bn), m.div_ceil(64)); let params = super::QsaGemmParams { matrix: [ m as i32, n as i32, k as i32, lda, ldb, n as i32, tn as i32, tm as i32, ], batch_strides: [ *a_batch.last().unwrap(), *b_batch.last().unwrap(), i64::from(m) * i64::from(n), ], swizzle, k_iterations: (k / bk) as i32, batch_ndim, padding: 0, }; let mut bindings = vec![ai(0), bi(1), output(3), bytes(4, ¶ms)]; let batch_strides = a_batch.iter().chain(&b_batch).copied().collect::>(); if batch_ndim > 1 { bindings.extend([ bytes(6, batch_shape.as_slice()), bytes(7, batch_strides.as_slice()), ]); } dispatch_specialized( &format!( "steel_gemm_fused_{}{trans}_{dtype_name}_{dtype_name}_bm64_bn{bn}_bk{bk}_wm2_wn{wn}", if nax { "nax_" } else { "" } ), &bindings, &constants(&[ (10, batch_ndim > 1), (100, false), (110, false), (200, m.is_multiple_of(64)), (201, n.is_multiple_of(bn)), (202, k.is_multiple_of(bk)), ]), [tn * (1 << swizzle), tm.div_ceil(1 << swizzle), batch], [32, wn, 2], )?; } } for copy in copies { encoder::add_temporaries(std::slice::from_ref(&*copy.buffer())); } Ok(()) } #[test] #[ignore = "requires Apple M5 Max; actual BF16 batched matmul reference"] fn mtplx_dense_attention_batches_match_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(); let bf16 = |shape: &[i32], salt| { Array::new( shape, Dtype::BF16, pattern(shape.iter().map(|&d| d as u32).product(), salt), ) .unwrap() }; let swap = |a: &Array| { let mut axes = (0..a.layout().shape().len() as i32).collect::>(); let rank = axes.len(); axes.swap(rank - 2, rank - 1); ops::transpose(a, &axes, gpu).unwrap() }; let cases = include_str!("../../../../tests/fixtures/mtplx-dense-batch.jsonl") .lines() .map(|s| serde_json::from_str::(s).unwrap()) .collect::>(); assert_eq!(cases.len(), 34); let mut checked = 0; for deferred in [false, true] { for case in &cases { super::eval::execution_tests::observed(|| { let kind = case["kind"].as_str().unwrap(); let rows = case["rows"].as_i64().unwrap() as i32; let (k, n) = (64, 37); let (a, b) = match kind { "qk" | "pv" => { let batch = case["batch"].as_i64().unwrap() as i32; let total = case["total"].as_i64().unwrap() as i32; if kind == "qk" { let q = ops::transpose( &bf16(&[batch, rows, 24, 256], 91), &[0, 2, 1, 3], gpu, ) .unwrap(); ( ops::reshape(&q, &[batch, 2, 12, rows, 256], gpu).unwrap(), swap(&bf16(&[batch, 2, 1, total, 256], 92)), ) } else { ( bf16(&[batch, 2, 12, rows, total], 91), bf16(&[batch, 2, 1, total, 256], 92), ) } } "cross" => (bf16(&[2, 1, rows, k], 91), swap(&bf16(&[1, 3, n, k], 92))), "transpose" => (swap(&bf16(&[3, k, rows], 91)), bf16(&[3, k, n], 92)), "copy" => ( ops::slice( &bf16(&[3, rows, k * 2], 91), &[0, 0, 1], &[3, rows, k * 2], &[1, 1, 2], gpu, ) .unwrap(), ops::slice( &bf16(&[3, k, n * 2], 92), &[0, 0, 1], &[3, k, n * 2], &[1, 1, 2], gpu, ) .unwrap(), ), "collapse" | "flatten" => { let b = swap(&bf16(&[n, k], 92)); ( bf16(&[3, rows, k], 91), if kind == "collapse" { ops::expand_dims(&b, &[0], gpu).unwrap() } else { b }, ) } "offset" => { let a = ops::slice( &bf16(&[3 * rows * k + 1], 91), &[1], &[3 * rows * k + 1], &[1], gpu, ) .unwrap(); let b = ops::slice( &bf16(&[3 * n * k + 1], 92), &[1], &[3 * n * k + 1], &[1], gpu, ) .unwrap(); ( ops::reshape(&a, &[3, rows, k], gpu).unwrap(), swap(&ops::reshape(&b, &[3, n, k], gpu).unwrap()), ) } _ => panic!("unknown matmul case"), }; if !deferred { ops::evaluate(&streams, &[a.clone(), b.clone()], gpu, false).unwrap(); } let output = matmul(&a, &b, gpu).unwrap(); drop((a, b)); let (_, calls) = capture_dispatches(|| { ops::evaluate(&streams, std::slice::from_ref(&output), gpu, false).unwrap() }); if matches!(kind, "qk" | "pv" | "cross") { assert!( !calls.iter().any(|r| r.0.starts_with("steel_gemm_splitk_")), "split-K is disabled for true batches" ); } verify( case["kernel"].as_str().unwrap(), &[(&output.buffer(), output.layout().size() as u32)], ); checked += 1; }); } } assert_eq!(checked, 68); streams.clear_streams().unwrap(); } #[test] #[ignore = "requires Apple M5 Max; canonical QSA F32 graph with the original TF32 setting"] fn mtplx_qsa_score_array_graph_matches_reference() { use super::super::{ gpu::Context, qwen_mtplx_tests::{pattern, verify_typed}, }; 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(); let mut checked = 0; for row in include_str!("../../../../tests/fixtures/mtplx-custom-kernels.jsonl") .lines() .map(|s| serde_json::from_str::(s).unwrap()) .filter(|r| { r["kernel"].as_str().unwrap().starts_with("qsa_f32_score_") && r["tf32"].as_bool() == Some(tf32_enabled()) }) { super::eval::execution_tests::observed(|| { let n = |k: &str| row[k].as_u64().unwrap() as u32; let backing = |count: u32, salt: u32, bf16: bool, shape: &[i32], strides: &[i64]| { let (dtype, buffer) = if bf16 { (Dtype::BF16, pattern(count, salt)) } else { let data = (0..count) .flat_map(|i| { let base = ((i * 17 + salt * 13) % 257) as i32 - 128; (base as f32 / 128.0 + (i % 7) as f32 / 65536.0).to_le_bytes() }) .collect::>(); (Dtype::F32, super::scalar_buffer(&data).unwrap()) }; let a = Array::unallocated(shape, dtype).unwrap(); let (no_broadcast, rc, cc) = a.layout().contiguity(strides).unwrap(); let flags = super::array::Flags { contiguous: no_broadcast == count as usize, row_contiguous: rc, col_contiguous: cc, }; a.set_strided_data(buffer, count as usize, strides, flags, 0) .unwrap(); a }; let strides = |key: &str| { row[key] .as_array() .unwrap() .iter() .map(|s| s.as_i64().unwrap()) .collect::>() }; let mut qs = vec![0]; qs.extend(strides("q_strides")); let mut ps = vec![0, 0]; ps.extend(strides("pooled_strides")); let q = backing( n("q_count"), 94, row["bf16"].as_bool().unwrap(), &[1, n("rows") as i32, 4, 128], &qs, ); let p = backing( n("pooled_count"), 95, false, &[1, 1, 128, n("blocks") as i32], &ps, ); let out = super::attention::qsa_scores(&q, &p, 128, gpu).unwrap(); drop((q, p)); ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap(); verify_typed( row["kernel"].as_str().unwrap(), &[(&out.buffer(), n("rows") * n("blocks"), 4)], ); checked += 1; }); } assert_eq!(checked, 192); } #[test] #[ignore = "requires Apple M5 Max; floating Matmul factory and original dispatch"] fn mtplx_dense_float_graph_matches_reference() { use super::super::{ gpu::Context, qwen_mtplx_tests::{capture_dispatches, pattern, verify_typed}, }; 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(); let dtype = |tag: &str| match tag { "bfloat16" => Dtype::BF16, "float16" => Dtype::F16, "float32" => Dtype::F32, _ => panic!("dtype"), }; let cases = include_str!("../../../../tests/fixtures/mtplx-dense-float.jsonl") .lines() .map(|s| serde_json::from_str::(s).unwrap()) .filter(|r| r["tf32"].as_bool() == Some(tf32_enabled())) .collect::>(); assert_eq!(cases.len(), 114); for prepared in [false, true] { for case in &cases { super::eval::execution_tests::observed(|| { let [m, n, k] = ["m", "n", "k"].map(|s| case[s].as_i64().unwrap() as i32); let kind = case["kind"].as_str().unwrap(); let (ta, tb) = (case["ta"].as_bool().unwrap(), case["tb"].as_bool().unwrap()); let source = |shape: &[i32], salt: u32, typ: Dtype| { let mut physical = shape.to_vec(); if kind == "copy" { *physical.last_mut().unwrap() *= 2; } let count = physical.iter().product::() + i32::from(kind == "offset"); let mut a = ops::astype( &Array::new(&[count], Dtype::BF16, pattern(count as u32, salt)).unwrap(), typ, false, gpu, ) .unwrap(); if typ == Dtype::F32 { let data = (0..count) .flat_map(|i| ((i % 7) as f32 / 65536.).to_le_bytes()) .collect::>(); let perturb = Array::new(&[count], Dtype::F32, super::scalar_buffer(&data).unwrap()) .unwrap(); a = super::binary::binary(&a, &perturb, super::binary::Binary::Add, gpu) .unwrap(); } if kind == "offset" { a = ops::slice(&a, &[1], &[count], &[1], gpu).unwrap(); } a = ops::reshape(&a, &physical, gpu).unwrap(); if kind == "copy" { let mut steps = vec![1; physical.len()]; *steps.last_mut().unwrap() = 2; a = ops::slice(&a, &vec![0; physical.len()], &physical, &steps, gpu) .unwrap(); } a }; let mut ashape = if kind == "cross" { vec![2, 1] } else if ["copy", "offset", "collapse", "flatten"].contains(&kind) { vec![3] } else { vec![] }; let mut bshape = if kind == "cross" { vec![1, 3] } else if ["copy", "offset"].contains(&kind) { vec![3] } else if kind == "collapse" { vec![1] } else { vec![] }; ashape.extend(if ta { [k, m] } else { [m, k] }); bshape.extend(if tb { [n, k] } else { [k, n] }); let mut a = source(&ashape, 171, dtype(case["a_dtype"].as_str().unwrap())); let mut b = source(&bshape, 172, dtype(case["b_dtype"].as_str().unwrap())); let transpose = |a: &Array| { let mut axes = (0..a.layout().shape().len() as i32).collect::>(); let n = axes.len(); axes.swap(n - 2, n - 1); ops::transpose(a, &axes, gpu).unwrap() }; if ta { a = transpose(&a); } if tb { b = transpose(&b); } if ["left-vector", "vectors"].contains(&kind) { a = ops::reshape(&a, &[k], gpu).unwrap(); } if ["right-vector", "vectors"].contains(&kind) { b = ops::reshape(&b, &[k], gpu).unwrap(); } if prepared { ops::evaluate(&streams, &[a.clone(), b.clone()], gpu, false).unwrap(); } let out = matmul(&a, &b, gpu).unwrap(); drop((a, b)); let shape = case["output_shape"] .as_array() .unwrap() .iter() .map(|d| d.as_i64().unwrap() as i32) .collect::>(); assert_eq!(out.layout().shape(), shape); let (_, calls) = capture_dispatches(|| { ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap() }); if out.layout().dtype() == Dtype::F32 { assert!(!calls.iter().any(|c| c.0.starts_with("gemv_wide_"))); if !tf32_enabled() { assert!(!calls.iter().any(|c| c.0.contains("_nax_"))); } let contract = match (case["profile"].as_u64().unwrap(), tf32_enabled()) { (4, true) => Some(( "steel_gemm_splitk_nax_nn_float32_float32_bm64_bn64_bk256_wm2_wn2", [3, 1, 1], [32, 2, 2], )), (4, false) => Some(( "steel_gemm_splitk_nn_float32_float32_bm16_bn16_bk16_wm2_wn2_MN_naligned_K_naligned", [2, 1, 8], [32, 2, 2], )), (12, true) => Some(( "steel_gemm_splitk_nax_nn_float32_float32_bm64_bn64_bk256_wm2_wn2", [2, 1, 1], [32, 2, 2], )), (12, false) => Some(( "steel_gemm_splitk_nn_float32_float32_bm32_bn32_bk16_wm2_wn2_MN_naligned_K_taligned", [2, 2, 2], [32, 2, 2], )), (24, true) => Some(( "steel_gemm_fused_nax_nt_float32_float32_bm64_bn128_bk256_wm2_wn4", [4, 1, 6], [32, 4, 2], )), (24, false) => Some(( "steel_gemm_fused_nt_float32_float32_bm64_bn64_bk16_wm2_wn2", [1, 1, 6], [32, 2, 2], )), (19, _) => Some(( "gemv_float32_bm4_bn1_sm1_sn32_tm4_tn4_nc0_axpby0", [2, 1, 1], [32, 1, 4], )), _ => None, }; if let Some((name, grid, threads)) = contract { assert_eq!( calls .iter() .filter(|c| c.0 == name && c.1 == grid && c.2 == threads && !c.3) .count(), 1, "{}: {calls:?}", case["kernel"] ); } } if kind == "vectors" { assert!( calls .iter() .any(|c| c.0.starts_with("dot_product_") && c.2 == [512, 1, 1] && c.3) ); } let out = ops::astype(&out, Dtype::F32, false, gpu).unwrap(); ops::evaluate(&streams, std::slice::from_ref(&out), gpu, false).unwrap(); verify_typed( case["kernel"].as_str().unwrap(), &[(&out.buffer(), out.layout().size() as u32, 4)], ); }); } } }