#![allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
use mlx_native::ops::ssm_conv::{
dispatch_ssm_conv, dispatch_ssm_conv_with_capture, SsmConvParams,
};
use mlx_native::{DType, KernelRegistry, MlxBuffer, MlxDevice};
fn setup() -> (MlxDevice, KernelRegistry) {
let device = MlxDevice::new().expect("MlxDevice::new");
let registry = KernelRegistry::new();
(device, registry)
}
fn upload_f32(device: &MlxDevice, data: &[f32]) -> MlxBuffer {
let mut buf = device
.alloc_buffer(data.len() * 4, DType::F32, vec![data.len()])
.expect("alloc");
buf.as_mut_slice::<f32>().expect("mut").copy_from_slice(data);
buf
}
fn rand_vec(seed: &mut u32, n: usize, scale: f32) -> Vec<f32> {
(0..n)
.map(|_| {
*seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
let r = (*seed >> 8) as f32 / ((1u32 << 24) as f32);
(r * 2.0 - 1.0) * scale
})
.collect()
}
fn build_params_buf(device: &MlxDevice, p: SsmConvParams) -> MlxBuffer {
let raw = [p.channels, p.n_tokens, p.n_seqs, p.k_width];
let mut buf = device
.alloc_buffer(16, DType::U32, vec![4])
.expect("params buf");
let dst = buf.as_mut_slice::<u32>().expect("params mut");
dst.copy_from_slice(&raw);
buf
}
fn run_legacy(
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &[f32],
kernel_w: &[f32],
state: &[f32],
p: SsmConvParams,
) -> (Vec<f32>, Vec<f32>) {
let x_buf = upload_f32(device, x);
let kw_buf = upload_f32(device, kernel_w);
let old_state = upload_f32(device, state);
let x_elems = (p.channels * p.n_tokens * p.n_seqs) as usize;
let s_elems = ((p.k_width - 1) * p.channels * p.n_seqs) as usize;
let y_buf = device
.alloc_buffer(x_elems * 4, DType::F32, vec![x_elems])
.expect("y alloc");
let new_state_buf = device
.alloc_buffer(s_elems * 4, DType::F32, vec![s_elems])
.expect("new state alloc");
let params_buf = build_params_buf(device, p);
let mut enc = device.command_encoder().expect("enc");
dispatch_ssm_conv(
&mut enc, registry, device.metal_device(),
&x_buf, &kw_buf, &old_state, &new_state_buf, &y_buf,
¶ms_buf, p,
)
.expect("dispatch ssm_conv");
enc.commit_and_wait().expect("commit");
(
y_buf.as_slice::<f32>().expect("y read").to_vec(),
new_state_buf.as_slice::<f32>().expect("new_state read").to_vec(),
)
}
fn run_capture(
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &[f32],
kernel_w: &[f32],
state: &[f32],
p: SsmConvParams,
) -> (Vec<f32>, Vec<f32>) {
let x_buf = upload_f32(device, x);
let kw_buf = upload_f32(device, kernel_w);
let old_state = upload_f32(device, state);
let x_elems = (p.channels * p.n_tokens * p.n_seqs) as usize;
let capture_elems = (p.n_seqs as usize)
* (p.n_tokens as usize)
* ((p.k_width - 1) as usize)
* (p.channels as usize);
let y_buf = device
.alloc_buffer(x_elems * 4, DType::F32, vec![x_elems])
.expect("y alloc");
let cap_buf = device
.alloc_buffer(capture_elems * 4, DType::F32, vec![capture_elems])
.expect("cap alloc");
let params_buf = build_params_buf(device, p);
let mut enc = device.command_encoder().expect("enc");
dispatch_ssm_conv_with_capture(
&mut enc, registry, device.metal_device(),
&x_buf, &kw_buf, &old_state, &y_buf, &cap_buf,
¶ms_buf, p,
)
.expect("dispatch ssm_conv_with_capture");
enc.commit_and_wait().expect("commit");
(
y_buf.as_slice::<f32>().expect("y read").to_vec(),
cap_buf.as_slice::<f32>().expect("cap read").to_vec(),
)
}
fn assert_byte_identical(label: &str, a: &[f32], b: &[f32]) {
assert_eq!(a.len(), b.len(), "{label}: len mismatch");
for (i, (&x, &y)) in a.iter().zip(b.iter()).enumerate() {
assert_eq!(
x.to_bits(),
y.to_bits(),
"{label}: byte mismatch at idx {} ({} vs {})",
i, x, y
);
}
}
#[test]
fn capture_y_byte_identical_qwen35_shape() {
let (device, mut registry) = setup();
let p = SsmConvParams {
channels: 8192,
n_tokens: 4,
n_seqs: 1,
k_width: 4,
};
let x_n = (p.channels * p.n_tokens * p.n_seqs) as usize;
let w_n = (p.k_width * p.channels) as usize;
let s_n = ((p.k_width - 1) * p.channels * p.n_seqs) as usize;
let mut seed = 0xCAFE;
let x = rand_vec(&mut seed, x_n, 0.1);
let w = rand_vec(&mut seed, w_n, 0.05);
let s = rand_vec(&mut seed, s_n, 0.05);
let (legacy_y, legacy_state) = run_legacy(&device, &mut registry, &x, &w, &s, p);
let (cap_y, cap_capture) = run_capture(&device, &mut registry, &x, &w, &s, p);
assert_byte_identical("qwen35 y", &cap_y, &legacy_y);
let per_t = ((p.k_width - 1) * p.channels) as usize;
let last_t_offset = ((p.n_tokens - 1) as usize) * per_t;
let cap_last_t = &cap_capture[last_t_offset..last_t_offset + per_t];
let k_minus1 = (p.k_width - 1) as usize;
let channels = p.channels as usize;
for i in 0..k_minus1 {
for c in 0..channels {
let legacy_idx = c * k_minus1 + i; let cap_idx = i * channels + c; assert_eq!(
legacy_state[legacy_idx].to_bits(),
cap_last_t[cap_idx].to_bits(),
"qwen35 last-t capture vs legacy state at i={i} c={c}",
);
}
}
}
#[test]
fn capture_y_byte_identical_small_shape() {
let (device, mut registry) = setup();
let p = SsmConvParams {
channels: 32,
n_tokens: 4,
n_seqs: 1,
k_width: 4,
};
let x_n = (p.channels * p.n_tokens * p.n_seqs) as usize;
let w_n = (p.k_width * p.channels) as usize;
let s_n = ((p.k_width - 1) * p.channels * p.n_seqs) as usize;
let mut seed = 0x1234;
let x = rand_vec(&mut seed, x_n, 0.1);
let w = rand_vec(&mut seed, w_n, 0.05);
let s = rand_vec(&mut seed, s_n, 0.05);
let (legacy_y, legacy_state) = run_legacy(&device, &mut registry, &x, &w, &s, p);
let (cap_y, cap_capture) = run_capture(&device, &mut registry, &x, &w, &s, p);
assert_byte_identical("small y", &cap_y, &legacy_y);
let per_t = ((p.k_width - 1) * p.channels) as usize;
let last_t_offset = ((p.n_tokens - 1) as usize) * per_t;
let cap_last_t = &cap_capture[last_t_offset..last_t_offset + per_t];
let k_minus1 = (p.k_width - 1) as usize;
let channels = p.channels as usize;
for i in 0..k_minus1 {
for c in 0..channels {
let legacy_idx = c * k_minus1 + i;
let cap_idx = i * channels + c;
assert_eq!(
legacy_state[legacy_idx].to_bits(),
cap_last_t[cap_idx].to_bits(),
"small last-t capture vs legacy state at i={i} c={c}",
);
}
}
}
#[test]
fn capture_intermediate_t_matches_truncated_dispatch() {
let (device, mut registry) = setup();
let n_full = 4u32;
let p_full = SsmConvParams {
channels: 64,
n_tokens: n_full,
n_seqs: 1,
k_width: 4,
};
let x_n = (p_full.channels * n_full * p_full.n_seqs) as usize;
let w_n = (p_full.k_width * p_full.channels) as usize;
let s_n = ((p_full.k_width - 1) * p_full.channels * p_full.n_seqs) as usize;
let mut seed = 0x517E;
let x = rand_vec(&mut seed, x_n, 0.1);
let w = rand_vec(&mut seed, w_n, 0.05);
let s = rand_vec(&mut seed, s_n, 0.05);
let (_cap_y, cap_capture) =
run_capture(&device, &mut registry, &x, &w, &s, p_full);
let per_t = ((p_full.k_width - 1) * p_full.channels) as usize;
let k_minus1 = (p_full.k_width - 1) as usize;
let channels = p_full.channels as usize;
for t in 0..(n_full as usize - 1) {
let truncated_x_n = (p_full.channels as usize) * (t + 1) * (p_full.n_seqs as usize);
let truncated_x = x[..truncated_x_n].to_vec();
let p_trunc = SsmConvParams {
channels: p_full.channels,
n_tokens: (t + 1) as u32,
n_seqs: p_full.n_seqs,
k_width: p_full.k_width,
};
let (_y_trunc, state_trunc) =
run_legacy(&device, &mut registry, &truncated_x, &w, &s, p_trunc);
let cap_t_offset = t * per_t;
let cap_t = &cap_capture[cap_t_offset..cap_t_offset + per_t];
for i in 0..k_minus1 {
for c in 0..channels {
let cap_idx = i * channels + c;
let legacy_idx = c * k_minus1 + i;
assert_eq!(
cap_t[cap_idx].to_bits(),
state_trunc[legacy_idx].to_bits(),
"capture[t={t}] vs trunc(n_tokens={trunc_n}) at i={i} c={c}",
trunc_n = t + 1,
);
}
}
}
}
#[test]
fn capture_intermediate_positions_differ() {
let (device, mut registry) = setup();
let p = SsmConvParams {
channels: 64,
n_tokens: 4,
n_seqs: 1,
k_width: 4,
};
let x_n = (p.channels * p.n_tokens * p.n_seqs) as usize;
let w_n = (p.k_width * p.channels) as usize;
let s_n = ((p.k_width - 1) * p.channels * p.n_seqs) as usize;
let mut seed = 0xDEAD;
let x = rand_vec(&mut seed, x_n, 0.5); let w = rand_vec(&mut seed, w_n, 0.5);
let s = rand_vec(&mut seed, s_n, 0.5);
let (_cap_y, cap_capture) = run_capture(&device, &mut registry, &x, &w, &s, p);
let per_t = ((p.k_width - 1) * p.channels) as usize;
let first_t = &cap_capture[0..per_t];
let last_t_offset = ((p.n_tokens - 1) as usize) * per_t;
let last_t = &cap_capture[last_t_offset..last_t_offset + per_t];
let differs = first_t
.iter()
.zip(last_t.iter())
.any(|(&a, &b)| a.to_bits() != b.to_bits());
assert!(
differs,
"capture[0] unexpectedly equals capture[last] — recurrence may be degenerate"
);
}