Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,970 @@
|
||||
//! 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, ©)?;
|
||||
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, ¶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::<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)],
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user