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

971 lines
36 KiB
Rust

//! 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<Array, String> {
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<Array, String> {
// 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<Array>,
) -> 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, &copy)?;
copies.push(copy.clone());
Ok((false, cols, copy))
}
pub(super) fn tf32_enabled() -> bool {
static ENABLED: std::sync::OnceLock<bool> = 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<i64>; 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::<Vec<_>>()
};
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, &params),
],
&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, &params)];
let batch_strides = a_batch.iter().chain(&b_batch).copied().collect::<Vec<_>>();
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::<Vec<_>>();
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::<serde_json::Value>(s).unwrap())
.collect::<Vec<_>>();
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::<serde_json::Value>(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::<Vec<_>>();
(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::<Vec<_>>()
};
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::<serde_json::Value>(s).unwrap())
.filter(|r| r["tf32"].as_bool() == Some(tf32_enabled()))
.collect::<Vec<_>>();
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>() + 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::<Vec<_>>();
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::<Vec<_>>();
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::<Vec<_>>();
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)],
);
});
}
}
}