971 lines
36 KiB
Rust
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, ©)?;
|
|
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)],
|
|
);
|
|
});
|
|
}
|
|
}
|
|
}
|