use mlx_native::ops::gated_delta_net::{cpu_reference_f32 as gdn_cpu_ref, GatedDeltaNetParams};
use crate::inference::models::qwen35::Qwen35Config;
#[derive(Debug, Clone)]
pub struct DeltaNetLayerWeights {
pub attn_norm: Vec<f32>,
pub post_attn_norm: Vec<f32>,
pub attn_qkv: Vec<f32>,
pub attn_gate: Vec<f32>,
pub ssm_conv1d: Vec<f32>,
pub ssm_alpha: Vec<f32>,
pub ssm_dt_bias: Vec<f32>,
pub ssm_beta: Vec<f32>,
pub ssm_a: Vec<f32>,
pub ssm_norm: Vec<f32>,
pub ssm_out: Vec<f32>,
}
#[derive(Debug, Clone, Copy)]
pub struct DeltaNetLayerShape {
pub hidden_size: u32,
pub n_k_heads: u32,
pub n_v_heads: u32,
pub d_k: u32,
pub d_v: u32,
pub conv_kernel: u32, pub rms_norm_eps: f32,
}
impl DeltaNetLayerShape {
pub fn from_config(cfg: &Qwen35Config) -> Self {
Self {
hidden_size: cfg.hidden_size,
n_k_heads: cfg.linear_num_key_heads,
n_v_heads: cfg.linear_num_value_heads,
d_k: cfg.linear_key_head_dim,
d_v: cfg.linear_value_head_dim,
conv_kernel: cfg.linear_conv_kernel_dim,
rms_norm_eps: cfg.rms_norm_eps,
}
}
pub fn qkv_channels(&self) -> u32 {
2 * self.n_k_heads * self.d_k + self.n_v_heads * self.d_v
}
}
fn rms_norm_row(x: &[f32], weight: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let sum_sq: f32 = x.iter().map(|v| v * v).sum();
let inv = (sum_sq / n + eps).sqrt().recip();
x.iter()
.zip(weight.iter())
.map(|(xi, wi)| xi * inv * wi)
.collect()
}
fn l2_norm_row(x: &[f32], eps: f32) -> Vec<f32> {
let sum_sq: f32 = x.iter().map(|v| v * v).sum();
let inv = (sum_sq + eps).sqrt().recip();
x.iter().map(|v| v * inv).collect()
}
fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn softplus(x: f32) -> f32 {
if x > 20.0 {
x
} else if x < -20.0 {
0.0
} else {
(1.0 + x.exp()).ln()
}
}
fn matmul_a_by_bt(lhs: &[f32], rhs: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for kk in 0..k {
acc += lhs[i * k + kk] * rhs[j * k + kk];
}
out[i * n + j] = acc;
}
}
out
}
fn ssm_conv_scalar(
x: &[f32], kernel: &[f32], conv_state: &[f32], seq: usize,
channels: usize,
k_width: usize,
) -> Vec<f32> {
let km1 = k_width - 1;
let mut out = vec![0.0f32; seq * channels];
for t in 0..seq {
for c in 0..channels {
let mut acc = 0.0f32;
for kk in 0..k_width {
let t_ext = t + kk;
let val = if t_ext < km1 {
conv_state[t_ext * channels + c]
} else {
x[(t_ext - km1) * channels + c]
};
acc += kernel[kk * channels + c] * val;
}
out[t * channels + c] = silu(acc);
}
}
out
}
pub fn delta_net_layer_cpu_ref(
x: &[f32],
weights: &DeltaNetLayerWeights,
shape: DeltaNetLayerShape,
state_in: &[f32],
conv_state: &[f32],
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let h = shape.hidden_size as usize;
let nk = shape.n_k_heads as usize;
let nv = shape.n_v_heads as usize;
let dk = shape.d_k as usize;
let dv = shape.d_v as usize;
let k_width = shape.conv_kernel as usize;
let km1 = k_width - 1;
let qkv_channels = shape.qkv_channels() as usize;
let z_channels = nv * dv;
let seq = x.len() / h;
assert_eq!(x.len(), seq * h);
assert_eq!(weights.attn_norm.len(), h);
assert_eq!(weights.attn_qkv.len(), qkv_channels * h);
assert_eq!(weights.attn_gate.len(), z_channels * h);
assert_eq!(weights.ssm_conv1d.len(), k_width * qkv_channels);
assert_eq!(weights.ssm_alpha.len(), nv * h);
assert_eq!(weights.ssm_dt_bias.len(), nv);
assert_eq!(weights.ssm_beta.len(), nv * h);
assert_eq!(weights.ssm_a.len(), nv);
assert_eq!(weights.ssm_norm.len(), dv);
assert_eq!(weights.ssm_out.len(), h * z_channels);
assert_eq!(state_in.len(), dk * dv * nv);
assert_eq!(conv_state.len(), km1 * qkv_channels);
let mut x_norm = vec![0.0f32; seq * h];
for t in 0..seq {
let row = &x[t * h..(t + 1) * h];
let normed = rms_norm_row(row, &weights.attn_norm, shape.rms_norm_eps);
x_norm[t * h..(t + 1) * h].copy_from_slice(&normed);
}
let qkv = matmul_a_by_bt(&x_norm, &weights.attn_qkv, seq, h, qkv_channels);
let z = matmul_a_by_bt(&x_norm, &weights.attn_gate, seq, h, z_channels);
let qkv_conv = ssm_conv_scalar(
&qkv,
&weights.ssm_conv1d,
conv_state,
seq,
qkv_channels,
k_width,
);
let q_span = nk * dk;
let k_span = nk * dk;
let v_span = nv * dv;
let mut q_buf = vec![0.0f32; seq * q_span];
let mut k_buf = vec![0.0f32; seq * k_span];
let mut v_buf = vec![0.0f32; seq * v_span];
for t in 0..seq {
let base = t * qkv_channels;
q_buf[t * q_span..(t + 1) * q_span].copy_from_slice(&qkv_conv[base..base + q_span]);
k_buf[t * k_span..(t + 1) * k_span]
.copy_from_slice(&qkv_conv[base + q_span..base + q_span + k_span]);
v_buf[t * v_span..(t + 1) * v_span]
.copy_from_slice(&qkv_conv[base + q_span + k_span..base + qkv_channels]);
}
let q_scale = 1.0 / (dk as f32).sqrt();
for t in 0..seq {
for h_idx in 0..nk {
let off = (t * nk + h_idx) * dk;
let row = &q_buf[off..off + dk];
let normed = l2_norm_row(row, shape.rms_norm_eps);
for (dst, v) in q_buf[off..off + dk].iter_mut().zip(normed.iter()) {
*dst = v * q_scale;
}
}
for h_idx in 0..nk {
let off = (t * nk + h_idx) * dk;
let row = &k_buf[off..off + dk];
let normed = l2_norm_row(row, shape.rms_norm_eps);
k_buf[off..off + dk].copy_from_slice(&normed);
}
}
let alpha_logits = matmul_a_by_bt(&x_norm, &weights.ssm_alpha, seq, h, nv);
let beta_logits = matmul_a_by_bt(&x_norm, &weights.ssm_beta, seq, h, nv);
let mut g = vec![0.0f32; seq * nv];
let mut beta = vec![0.0f32; seq * nv];
for t in 0..seq {
for h_idx in 0..nv {
let a_logit = alpha_logits[t * nv + h_idx] + weights.ssm_dt_bias[h_idx];
g[t * nv + h_idx] = softplus(a_logit) * (-weights.ssm_a[h_idx]);
beta[t * nv + h_idx] = sigmoid(beta_logits[t * nv + h_idx]);
}
}
let q_trans = transpose_for_gdn(&q_buf, seq, nk, dk);
let k_trans = transpose_for_gdn(&k_buf, seq, nk, dk);
let v_trans = transpose_for_gdn(&v_buf, seq, nv, dv);
let params = GatedDeltaNetParams {
d_k: dk as u32,
d_v: dv as u32,
n_k_heads: nk as u32,
n_v_heads: nv as u32,
n_tokens: seq as u32,
n_seqs: 1,
};
let (gdn_out_mlx, new_state) =
gdn_cpu_ref(&q_trans, &k_trans, &v_trans, &g, &beta, state_in, params);
let mut attn_out = vec![0.0f32; seq * nv * dv];
for t in 0..seq {
for vh in 0..nv {
for d in 0..dv {
let src = t * nv * dv + vh * dv + d;
attn_out[t * nv * dv + vh * dv + d] = gdn_out_mlx[src];
}
}
}
assert_eq!(
weights.ssm_norm.len(),
dv,
"ssm_norm shape mismatch: expected [D_v={}] got {}",
dv,
weights.ssm_norm.len()
);
let mut gated = vec![0.0f32; seq * z_channels];
for t in 0..seq {
for vh in 0..nv {
let head_off = t * z_channels + vh * dv;
let head_row = &attn_out[head_off..head_off + dv];
let normed = rms_norm_row(head_row, &weights.ssm_norm, shape.rms_norm_eps);
for d in 0..dv {
let z_val = z[head_off + d];
let z_silu = z_val / (1.0 + (-z_val).exp());
gated[head_off + d] = normed[d] * z_silu;
}
}
}
let output = matmul_a_by_bt(&gated, &weights.ssm_out, seq, z_channels, h);
let mut new_conv_state = vec![0.0f32; km1 * qkv_channels];
for i in 0..km1 {
let t_ext = seq + i; if t_ext < km1 {
new_conv_state[i * qkv_channels..(i + 1) * qkv_channels]
.copy_from_slice(&conv_state[t_ext * qkv_channels..(t_ext + 1) * qkv_channels]);
} else {
let t_in = t_ext - km1;
new_conv_state[i * qkv_channels..(i + 1) * qkv_channels]
.copy_from_slice(&qkv[t_in * qkv_channels..(t_in + 1) * qkv_channels]);
}
}
(output, new_state, new_conv_state)
}
fn transpose_for_gdn(src: &[f32], seq: usize, n_heads: usize, d: usize) -> Vec<f32> {
let mut dst = vec![0.0f32; seq * n_heads * d];
dst.copy_from_slice(src);
dst
}
#[cfg(test)]
mod tests {
use super::*;
fn mk_rand(seed: &mut u32, n: usize, scale: f32) -> Vec<f32> {
(0..n)
.map(|_| {
*seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((*seed as i32 as f32) / (i32::MAX as f32)) * scale
})
.collect()
}
fn small_shape() -> DeltaNetLayerShape {
DeltaNetLayerShape {
hidden_size: 8,
n_k_heads: 2,
n_v_heads: 4, d_k: 4,
d_v: 4,
conv_kernel: 4,
rms_norm_eps: 1e-6,
}
}
fn synthetic_weights(shape: DeltaNetLayerShape, seed_init: u32) -> DeltaNetLayerWeights {
let h = shape.hidden_size as usize;
let _nk = shape.n_k_heads as usize;
let nv = shape.n_v_heads as usize;
let _dk = shape.d_k as usize;
let dv = shape.d_v as usize;
let k_width = shape.conv_kernel as usize;
let qkv_channels = shape.qkv_channels() as usize;
let z_channels = nv * dv;
let mut seed = seed_init;
DeltaNetLayerWeights {
attn_norm: {
let mut v = vec![1.0f32; h];
for (i, x) in v.iter_mut().enumerate() {
*x += 0.01 * (i as f32);
}
v
},
post_attn_norm: vec![1.0f32; h],
attn_qkv: mk_rand(&mut seed, qkv_channels * h, 0.1),
attn_gate: mk_rand(&mut seed, z_channels * h, 0.1),
ssm_conv1d: mk_rand(&mut seed, k_width * qkv_channels, 0.1),
ssm_alpha: mk_rand(&mut seed, nv * h, 0.1),
ssm_dt_bias: mk_rand(&mut seed, nv, 0.05),
ssm_beta: mk_rand(&mut seed, nv * h, 0.1),
ssm_a: mk_rand(&mut seed, nv, 0.1),
ssm_norm: {
let mut v = vec![1.0f32; dv];
for (i, x) in v.iter_mut().enumerate() {
*x += 0.01 * (i as f32);
}
v
},
ssm_out: mk_rand(&mut seed, h * z_channels, 0.1),
}
}
#[test]
fn shape_qkv_channels() {
let s = small_shape();
assert_eq!(s.qkv_channels(), 32);
}
#[test]
fn delta_net_layer_produces_expected_shape() {
let shape = small_shape();
let weights = synthetic_weights(shape, 0xABCD);
let seq_len = 3;
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let x: Vec<f32> = (0..seq_len * h).map(|i| 0.01 * i as f32).collect();
let state_in = vec![0.0f32; (shape.d_k * shape.d_v * shape.n_v_heads) as usize];
let conv_state = vec![0.0f32; km1 * qkv_channels];
let (out, new_state, new_conv) =
delta_net_layer_cpu_ref(&x, &weights, shape, &state_in, &conv_state);
assert_eq!(out.len(), seq_len * h);
assert_eq!(new_state.len(), state_in.len());
assert_eq!(new_conv.len(), conv_state.len());
assert!(out.iter().all(|v| v.is_finite()));
let sum_abs: f32 = out.iter().map(|v| v.abs()).sum();
assert!(sum_abs > 1e-6, "output is nearly zero — something broken");
}
#[test]
fn delta_net_layer_deterministic() {
let shape = small_shape();
let weights = synthetic_weights(shape, 0x1234);
let seq_len = 2;
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let x: Vec<f32> = (0..seq_len * h).map(|i| 0.02 * i as f32).collect();
let state_in = vec![0.0f32; (shape.d_k * shape.d_v * shape.n_v_heads) as usize];
let conv_state = vec![0.0f32; km1 * qkv_channels];
let (o1, s1, c1) = delta_net_layer_cpu_ref(&x, &weights, shape, &state_in, &conv_state);
let (o2, s2, c2) = delta_net_layer_cpu_ref(&x, &weights, shape, &state_in, &conv_state);
for i in 0..o1.len() {
assert_eq!(o1[i].to_bits(), o2[i].to_bits());
}
for i in 0..s1.len() {
assert_eq!(s1[i].to_bits(), s2[i].to_bits());
}
for i in 0..c1.len() {
assert_eq!(c1[i].to_bits(), c2[i].to_bits());
}
}
#[test]
fn delta_net_layer_state_in_affects_output() {
let shape = small_shape();
let weights = synthetic_weights(shape, 0x5678);
let seq_len = 2;
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let x: Vec<f32> = (0..seq_len * h).map(|i| 0.03 * i as f32).collect();
let conv_state = vec![0.0f32; km1 * qkv_channels];
let state_zeros = vec![0.0f32; (shape.d_k * shape.d_v * shape.n_v_heads) as usize];
let mut state_nonzero = state_zeros.clone();
for (i, v) in state_nonzero.iter_mut().enumerate() {
*v = 0.1 * ((i % 13) as f32);
}
let (o_zero, _, _) =
delta_net_layer_cpu_ref(&x, &weights, shape, &state_zeros, &conv_state);
let (o_nonzero, _, _) =
delta_net_layer_cpu_ref(&x, &weights, shape, &state_nonzero, &conv_state);
let mut any_diff = false;
for i in 0..o_zero.len() {
if (o_zero[i] - o_nonzero[i]).abs() > 1e-5 {
any_diff = true;
break;
}
}
assert!(any_diff, "initial state had no effect on output");
}
#[test]
fn delta_net_layer_chunked_equals_monolithic() {
let shape = small_shape();
let weights = synthetic_weights(shape, 0xFACE);
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let state_zeros = vec![0.0f32; (shape.d_k * shape.d_v * shape.n_v_heads) as usize];
let conv_zeros = vec![0.0f32; km1 * qkv_channels];
let x_full: Vec<f32> = (0..2 * h).map(|i| 0.1 + 0.01 * i as f32).collect();
let (out_mono, _, _) =
delta_net_layer_cpu_ref(&x_full, &weights, shape, &state_zeros, &conv_zeros);
let x_t0 = x_full[0..h].to_vec();
let x_t1 = x_full[h..2 * h].to_vec();
let (out_t0, state_after_t0, conv_after_t0) =
delta_net_layer_cpu_ref(&x_t0, &weights, shape, &state_zeros, &conv_zeros);
let (out_t1, _, _) =
delta_net_layer_cpu_ref(&x_t1, &weights, shape, &state_after_t0, &conv_after_t0);
for i in 0..h {
let d = (out_mono[i] - out_t0[i]).abs();
assert!(
d < 1e-5,
"chunked-vs-mono t0 mismatch at {}: mono={}, chunk={}",
i,
out_mono[i],
out_t0[i]
);
}
for i in 0..h {
let d = (out_mono[h + i] - out_t1[i]).abs();
assert!(
d < 1e-5,
"chunked-vs-mono t1 mismatch at {}: mono={}, chunk={}",
i,
out_mono[h + i],
out_t1[i]
);
}
}
#[test]
fn delta_net_layer_rejects_wrong_state_shape() {
let shape = small_shape();
let weights = synthetic_weights(shape, 0xDEAD);
let seq_len = 1;
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let x: Vec<f32> = (0..seq_len * h).map(|i| 0.01 * i as f32).collect();
let wrong_state = vec![0.0f32; 123];
let conv_state = vec![0.0f32; km1 * qkv_channels];
let res = std::panic::catch_unwind(|| {
delta_net_layer_cpu_ref(&x, &weights, shape, &wrong_state, &conv_state);
});
assert!(res.is_err(), "wrong state shape should panic");
}
}