use crate::attention::gdn::{sigmoid, softplus};
use crate::backward::ops::lora_vjp;
use crate::model::qwen35_config::Qwen35Config;
pub struct GdnGrads {
pub grad_a_qkv: Vec<f32>,
pub grad_b_qkv: Vec<f32>,
pub grad_a_z: Vec<f32>,
pub grad_b_z: Vec<f32>,
pub grad_a_b: Vec<f32>,
pub grad_b_b: Vec<f32>,
pub grad_a_a: Vec<f32>,
pub grad_b_a: Vec<f32>,
pub grad_a_out: Vec<f32>,
pub grad_b_out: Vec<f32>,
pub dx: Vec<f32>,
}
pub struct GdnSaved {
pub seq_len: usize,
pub num_key_heads: usize,
pub num_value_heads: usize,
pub ratio: usize,
pub key_dim: usize,
pub value_dim: usize,
pub hidden_size: usize,
pub qkv_dim: usize,
pub output_dim: usize,
pub kernel_size: usize,
pub scale: f32,
pub rms_eps: f32,
pub inputs: Vec<f32>,
pub qkv_proj: Vec<f32>,
pub z_proj: Vec<f32>,
pub beta_raw: Vec<f32>,
pub alpha_proj: Vec<f32>,
pub beta: Vec<f32>,
pub g: Vec<f32>,
pub conv_out: Vec<f32>,
pub conv_buffers: Vec<f32>,
pub q_hat: Vec<f32>,
pub k_hat: Vec<f32>,
pub v: Vec<f32>,
pub q_norm: Vec<f32>,
pub k_norm: Vec<f32>,
pub q_eps_norm: Vec<f32>,
pub k_eps_norm: Vec<f32>,
pub kv_mem: Vec<f32>,
pub s_after: Vec<f32>,
pub o_heads: Vec<f32>,
pub rms_vals: Vec<f32>,
pub silu_z: Vec<f32>,
pub lora_rank: usize,
pub lora_scale: f32,
pub h_qkv: Vec<f32>,
pub h_z: Vec<f32>,
pub h_b: Vec<f32>,
pub h_a: Vec<f32>,
pub h_out: Vec<f32>,
pub gated_buf: Vec<f32>,
pub lora_a_qkv: Vec<f32>,
pub lora_b_qkv: Vec<f32>,
pub lora_a_z: Vec<f32>,
pub lora_b_z: Vec<f32>,
pub lora_a_b: Vec<f32>,
pub lora_b_b: Vec<f32>,
pub lora_a_a: Vec<f32>,
pub lora_b_a: Vec<f32>,
pub lora_a_out: Vec<f32>,
pub lora_b_out: Vec<f32>,
}
impl GdnSaved {
pub fn new(
seq_len: usize,
num_key_heads: usize,
value_heads: usize,
key_dim: usize,
value_dim: usize,
hidden_size: usize,
qkv_dim: usize,
output_dim: usize,
kernel_size: usize,
scale: f32,
rms_eps: f32,
) -> Self {
let ratio = if num_key_heads == 0 {
1
} else {
value_heads / num_key_heads
};
let buf_len = kernel_size.saturating_sub(1);
Self {
seq_len,
num_key_heads,
num_value_heads: value_heads,
ratio,
key_dim,
value_dim,
hidden_size,
qkv_dim,
output_dim,
kernel_size,
scale,
rms_eps,
inputs: vec![0.0; seq_len * hidden_size],
qkv_proj: vec![0.0; seq_len * qkv_dim],
z_proj: vec![0.0; seq_len * output_dim],
beta_raw: vec![0.0; seq_len * value_heads],
alpha_proj: vec![0.0; seq_len * value_heads],
beta: vec![0.0; seq_len * value_heads],
g: vec![0.0; seq_len * value_heads],
conv_out: vec![0.0; seq_len * qkv_dim],
conv_buffers: vec![0.0; seq_len * qkv_dim * buf_len],
q_hat: vec![0.0; seq_len * value_heads * key_dim],
k_hat: vec![0.0; seq_len * value_heads * key_dim],
v: vec![0.0; seq_len * value_heads * value_dim],
q_norm: vec![0.0; seq_len * value_heads],
k_norm: vec![0.0; seq_len * value_heads],
q_eps_norm: vec![0.0; seq_len * value_heads],
k_eps_norm: vec![0.0; seq_len * value_heads],
kv_mem: vec![0.0; seq_len * value_heads * value_dim],
s_after: vec![0.0; seq_len * value_heads * key_dim * value_dim],
o_heads: vec![0.0; seq_len * value_heads * value_dim],
rms_vals: vec![0.0; seq_len * value_heads],
silu_z: vec![0.0; seq_len * value_heads * value_dim],
lora_rank: 0,
lora_scale: 0.0,
h_qkv: Vec::new(),
h_z: Vec::new(),
h_b: Vec::new(),
h_a: Vec::new(),
h_out: Vec::new(),
gated_buf: Vec::new(),
lora_a_qkv: Vec::new(),
lora_b_qkv: Vec::new(),
lora_a_z: Vec::new(),
lora_b_z: Vec::new(),
lora_a_b: Vec::new(),
lora_b_b: Vec::new(),
lora_a_a: Vec::new(),
lora_b_a: Vec::new(),
lora_a_out: Vec::new(),
lora_b_out: Vec::new(),
}
}
}
#[allow(clippy::too_many_arguments)]
type LoraBound<'a> = (
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
&'a [f32],
);
pub fn gdn_forward_save(
inputs: &[f32],
weights: &crate::attention::gdn::GatedDeltaNetWeights,
_cfg: &Qwen35Config,
saved: &mut GdnSaved,
outputs: &mut [f32],
lora_a_qkv: Option<&[f32]>,
lora_b_qkv: Option<&[f32]>,
lora_a_z: Option<&[f32]>,
lora_b_z: Option<&[f32]>,
lora_a_b: Option<&[f32]>,
lora_b_b: Option<&[f32]>,
lora_a_a: Option<&[f32]>,
lora_b_a: Option<&[f32]>,
lora_a_out: Option<&[f32]>,
lora_b_out: Option<&[f32]>,
lora_rank: usize,
lora_scale: f32,
) {
use crate::forward::cpu::matmul_bt;
let seq_len = saved.seq_len;
let hidden = saved.hidden_size;
let num_kh = saved.num_key_heads;
let value_heads = saved.num_value_heads;
let ratio = saved.ratio;
let key_dim = saved.key_dim;
let value_dim = saved.value_dim;
let qkv_dim = saved.qkv_dim;
let output_dim = saved.output_dim;
let kernel_size = saved.kernel_size;
let scale = saved.scale;
let rms_eps = saved.rms_eps;
let buf_len = kernel_size.saturating_sub(1);
let q_total = num_kh * key_dim;
let lora_bound: Option<LoraBound> = if lora_rank > 0 {
match (
lora_a_qkv, lora_b_qkv, lora_a_z, lora_b_z, lora_a_b, lora_b_b, lora_a_a, lora_b_a,
lora_a_out, lora_b_out,
) {
(
Some(a_qkv),
Some(b_qkv),
Some(a_z),
Some(b_z),
Some(a_b),
Some(b_b),
Some(a_a),
Some(b_a),
Some(a_out),
Some(b_out),
) => Some((a_qkv, b_qkv, a_z, b_z, a_b, b_b, a_a, b_a, a_out, b_out)),
_ => None,
}
} else {
None
};
saved.lora_rank = 0;
saved.lora_scale = 0.0;
saved.h_qkv.clear();
saved.h_z.clear();
saved.h_b.clear();
saved.h_a.clear();
saved.h_out.clear();
saved.gated_buf.clear();
saved.lora_a_qkv.clear();
saved.lora_b_qkv.clear();
saved.lora_a_z.clear();
saved.lora_b_z.clear();
saved.lora_a_b.clear();
saved.lora_b_b.clear();
saved.lora_a_a.clear();
saved.lora_b_a.clear();
saved.lora_a_out.clear();
saved.lora_b_out.clear();
if let Some((a_qkv, b_qkv, a_z, b_z, a_b, b_b, a_a, b_a, a_out, b_out)) = lora_bound {
saved.lora_rank = lora_rank;
saved.lora_scale = lora_scale;
saved.h_qkv = vec![0.0f32; seq_len * lora_rank];
saved.h_z = vec![0.0f32; seq_len * lora_rank];
saved.h_b = vec![0.0f32; seq_len * lora_rank];
saved.h_a = vec![0.0f32; seq_len * lora_rank];
saved.h_out = vec![0.0f32; seq_len * lora_rank];
saved.gated_buf = vec![0.0f32; seq_len * output_dim];
saved.lora_a_qkv = a_qkv.to_vec();
saved.lora_b_qkv = b_qkv.to_vec();
saved.lora_a_z = a_z.to_vec();
saved.lora_b_z = b_z.to_vec();
saved.lora_a_b = a_b.to_vec();
saved.lora_b_b = b_b.to_vec();
saved.lora_a_a = a_a.to_vec();
saved.lora_b_a = b_a.to_vec();
saved.lora_a_out = a_out.to_vec();
saved.lora_b_out = b_out.to_vec();
}
let mut conv_buf_live = vec![0.0f32; qkv_dim * buf_len];
let mut s_live = vec![0.0f32; value_heads * key_dim * value_dim];
for t in 0..seq_len {
let x = &inputs[t * hidden..(t + 1) * hidden];
saved.inputs[t * hidden..(t + 1) * hidden].copy_from_slice(x);
let qkv_out = &mut saved.qkv_proj[t * qkv_dim..(t + 1) * qkv_dim];
matmul_bt(x, &weights.in_proj_qkv, qkv_out, 1, hidden, qkv_dim);
if let Some((a_qkv, b_qkv, _, _, _, _, _, _, _, _)) = lora_bound {
let h = &mut saved.h_qkv[t * lora_rank..(t + 1) * lora_rank];
for r in 0..lora_rank {
h[r] = a_qkv[r * hidden..(r + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, xi)| a * xi)
.sum();
}
let qkv_out = &mut saved.qkv_proj[t * qkv_dim..(t + 1) * qkv_dim];
for i in 0..qkv_dim {
let acc: f32 = lora_scale
* b_qkv[i * lora_rank..(i + 1) * lora_rank]
.iter()
.zip(h.iter())
.map(|(b, hi)| b * hi)
.sum::<f32>();
qkv_out[i] += acc;
}
}
let z_out = &mut saved.z_proj[t * output_dim..(t + 1) * output_dim];
matmul_bt(x, &weights.in_proj_z, z_out, 1, hidden, output_dim);
if let Some((_, _, a_z, b_z, _, _, _, _, _, _)) = lora_bound {
let h = &mut saved.h_z[t * lora_rank..(t + 1) * lora_rank];
for r in 0..lora_rank {
h[r] = a_z[r * hidden..(r + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, xi)| a * xi)
.sum();
}
let z_out = &mut saved.z_proj[t * output_dim..(t + 1) * output_dim];
for i in 0..output_dim {
let acc: f32 = lora_scale
* b_z[i * lora_rank..(i + 1) * lora_rank]
.iter()
.zip(h.iter())
.map(|(b, hi)| b * hi)
.sum::<f32>();
z_out[i] += acc;
}
}
let beta_out = &mut saved.beta_raw[t * value_heads..(t + 1) * value_heads];
matmul_bt(x, &weights.in_proj_b, beta_out, 1, hidden, value_heads);
if let Some((_, _, _, _, a_b, b_b, _, _, _, _)) = lora_bound {
let h = &mut saved.h_b[t * lora_rank..(t + 1) * lora_rank];
for r in 0..lora_rank {
h[r] = a_b[r * hidden..(r + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, xi)| a * xi)
.sum();
}
let beta_out = &mut saved.beta_raw[t * value_heads..(t + 1) * value_heads];
for i in 0..value_heads {
let acc: f32 = lora_scale
* b_b[i * lora_rank..(i + 1) * lora_rank]
.iter()
.zip(h.iter())
.map(|(b, hi)| b * hi)
.sum::<f32>();
beta_out[i] += acc;
}
}
let alpha_out = &mut saved.alpha_proj[t * value_heads..(t + 1) * value_heads];
matmul_bt(x, &weights.in_proj_a, alpha_out, 1, hidden, value_heads);
if let Some((_, _, _, _, _, _, a_a, b_a, _, _)) = lora_bound {
let h = &mut saved.h_a[t * lora_rank..(t + 1) * lora_rank];
for r in 0..lora_rank {
h[r] = a_a[r * hidden..(r + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, xi)| a * xi)
.sum();
}
let alpha_out = &mut saved.alpha_proj[t * value_heads..(t + 1) * value_heads];
for i in 0..value_heads {
let acc: f32 = lora_scale
* b_a[i * lora_rank..(i + 1) * lora_rank]
.iter()
.zip(h.iter())
.map(|(b, hi)| b * hi)
.sum::<f32>();
alpha_out[i] += acc;
}
}
for vh in 0..value_heads {
let raw = saved.beta_raw[t * value_heads + vh];
saved.beta[t * value_heads + vh] = sigmoid(raw);
}
for vh in 0..value_heads {
let alpha_h = saved.alpha_proj[t * value_heads + vh];
let a = weights.a_log[vh].exp();
let sp = softplus(alpha_h + weights.dt_bias[vh]);
saved.g[t * value_heads + vh] = (-a * sp).exp();
}
let cb_off = t * qkv_dim * buf_len;
saved.conv_buffers[cb_off..cb_off + qkv_dim * buf_len].copy_from_slice(&conv_buf_live);
let conv_out_t = &mut saved.conv_out[t * qkv_dim..(t + 1) * qkv_dim];
let qkv_in = &saved.qkv_proj[t * qkv_dim..(t + 1) * qkv_dim];
conv1d_silu_fwd(
qkv_in,
&mut conv_buf_live,
&weights.conv1d_weight,
conv_out_t,
qkv_dim,
kernel_size,
);
for h in 0..value_heads {
let kh = h / ratio;
let q_start = kh * key_dim;
let k_start = q_total + kh * key_dim;
let v_start = q_total * 2 + h * value_dim;
let mut q_raw =
saved.conv_out[t * qkv_dim + q_start..t * qkv_dim + q_start + key_dim].to_vec();
let mut k_raw =
saved.conv_out[t * qkv_dim + k_start..t * qkv_dim + k_start + key_dim].to_vec();
let v_slice = &saved.conv_out[t * qkv_dim + v_start..t * qkv_dim + v_start + value_dim];
let v_off = (t * value_heads + h) * value_dim;
saved.v[v_off..v_off + value_dim].copy_from_slice(v_slice);
let q_sum_sq = l2_norm_sq(&q_raw);
let k_sum_sq = l2_norm_sq(&k_raw);
let q_norm = q_sum_sq.sqrt().max(1e-6_f32.sqrt());
let k_norm = k_sum_sq.sqrt().max(1e-6_f32.sqrt());
saved.q_norm[t * value_heads + h] = q_norm;
saved.k_norm[t * value_heads + h] = k_norm;
let q_eps_norm = (q_sum_sq + 1e-6).sqrt();
let k_eps_norm = (k_sum_sq + 1e-6).sqrt();
saved.q_eps_norm[t * value_heads + h] = q_eps_norm;
saved.k_eps_norm[t * value_heads + h] = k_eps_norm;
for v in &mut q_raw {
*v /= q_eps_norm;
}
for v in &mut k_raw {
*v /= k_eps_norm;
}
let q_off = (t * value_heads + h) * key_dim;
let k_off = (t * value_heads + h) * key_dim;
saved.q_hat[q_off..q_off + key_dim].copy_from_slice(&q_raw);
saved.k_hat[k_off..k_off + key_dim].copy_from_slice(&k_raw);
let g_h = saved.g[t * value_heads + h];
let beta_h = saved.beta[t * value_heads + h];
let s_off = h * key_dim * value_dim;
let s = &s_live[s_off..s_off + key_dim * value_dim];
let kvm_off = (t * value_heads + h) * value_dim;
let kv_mem_h = &mut saved.kv_mem[kvm_off..kvm_off + value_dim];
kv_mem_h.fill(0.0);
for i in 0..key_dim {
let ki = k_raw[i];
for j in 0..value_dim {
kv_mem_h[j] += s[i * value_dim + j] * ki;
}
}
let mut delta = vec![0.0f32; value_dim];
for j in 0..value_dim {
delta[j] = (v_slice[j] - kv_mem_h[j] * g_h) * beta_h;
}
let s_mut = &mut s_live[s_off..s_off + key_dim * value_dim];
for i in 0..key_dim {
for j in 0..value_dim {
s_mut[i * value_dim + j] = s_mut[i * value_dim + j] * g_h + k_raw[i] * delta[j];
}
}
let sa_off = (t * value_heads + h) * key_dim * value_dim;
saved.s_after[sa_off..sa_off + key_dim * value_dim].copy_from_slice(s_mut);
let o_off = (t * value_heads + h) * value_dim;
let o_h = &mut saved.o_heads[o_off..o_off + value_dim];
o_h.fill(0.0);
for i in 0..key_dim {
let qi = q_raw[i];
for j in 0..value_dim {
o_h[j] += s_mut[i * value_dim + j] * qi;
}
}
for j in 0..value_dim {
o_h[j] *= scale;
}
}
let z_slice = &saved.z_proj[t * output_dim..(t + 1) * output_dim];
let gamma = &weights.norm_weight[..value_dim];
let mut gated_buf = vec![0.0f32; output_dim];
for h in 0..value_heads {
let o_off = (t * value_heads + h) * value_dim;
let o_h = &saved.o_heads[o_off..o_off + value_dim];
let z_h = &z_slice[h * value_dim..(h + 1) * value_dim];
let sum_sq: f32 = o_h.iter().map(|v| v * v).sum();
let rms = (sum_sq / value_dim as f32 + rms_eps).sqrt();
saved.rms_vals[t * value_heads + h] = rms;
let inv_rms = 1.0 / rms;
let sz_off = (t * value_heads + h) * value_dim;
for j in 0..value_dim {
let sz = silu_f32(z_h[j]);
saved.silu_z[sz_off + j] = sz;
gated_buf[h * value_dim + j] = (o_h[j] * inv_rms) * gamma[j] * sz;
}
}
if lora_bound.is_some() {
saved.gated_buf[t * output_dim..(t + 1) * output_dim].copy_from_slice(&gated_buf);
}
let y_t = &mut outputs[t * hidden..(t + 1) * hidden];
matmul_bt(&gated_buf, &weights.out_proj, y_t, 1, output_dim, hidden);
if let Some((_, _, _, _, _, _, _, _, a_out, b_out)) = lora_bound {
let h_slot = &mut saved.h_out[t * lora_rank..(t + 1) * lora_rank];
for r in 0..lora_rank {
h_slot[r] = a_out[r * output_dim..(r + 1) * output_dim]
.iter()
.zip(gated_buf.iter())
.map(|(a, xi)| a * xi)
.sum();
}
let y_t = &mut outputs[t * hidden..(t + 1) * hidden];
for i in 0..hidden {
let acc: f32 = lora_scale
* b_out[i * lora_rank..(i + 1) * lora_rank]
.iter()
.zip(h_slot.iter())
.map(|(b, hi)| b * hi)
.sum::<f32>();
y_t[i] += acc;
}
}
}
}
pub fn gdn_backward(
grad_outputs: &[f32],
saved: &GdnSaved,
weights: &crate::attention::gdn::GatedDeltaNetWeights,
) -> GdnGrads {
let seq_len = saved.seq_len;
let hidden = saved.hidden_size;
let num_kh = saved.num_key_heads;
let value_heads = saved.num_value_heads;
let ratio = saved.ratio;
let key_dim = saved.key_dim;
let value_dim = saved.value_dim;
let qkv_dim = saved.qkv_dim;
let output_dim = saved.output_dim;
let kernel_size = saved.kernel_size;
let scale = saved.scale;
let buf_len = kernel_size.saturating_sub(1);
let q_total = num_kh * key_dim;
let gamma = &weights.norm_weight[..value_dim];
let have_lora = saved.lora_rank > 0 && !saved.h_qkv.is_empty();
let lora_rank = saved.lora_rank;
let lora_scale = saved.lora_scale;
let mut grad_a_qkv = if have_lora {
vec![0.0f32; lora_rank * hidden]
} else {
Vec::new()
};
let mut grad_b_qkv = if have_lora {
vec![0.0f32; qkv_dim * lora_rank]
} else {
Vec::new()
};
let mut grad_a_z = if have_lora {
vec![0.0f32; lora_rank * hidden]
} else {
Vec::new()
};
let mut grad_b_z = if have_lora {
vec![0.0f32; output_dim * lora_rank]
} else {
Vec::new()
};
let mut grad_a_b = if have_lora {
vec![0.0f32; lora_rank * hidden]
} else {
Vec::new()
};
let mut grad_b_b = if have_lora {
vec![0.0f32; value_heads * lora_rank]
} else {
Vec::new()
};
let mut grad_a_a = if have_lora {
vec![0.0f32; lora_rank * hidden]
} else {
Vec::new()
};
let mut grad_b_a = if have_lora {
vec![0.0f32; value_heads * lora_rank]
} else {
Vec::new()
};
let mut grad_a_out = if have_lora {
vec![0.0f32; lora_rank * output_dim]
} else {
Vec::new()
};
let mut grad_b_out = if have_lora {
vec![0.0f32; hidden * lora_rank]
} else {
Vec::new()
};
let mut grad_inputs = vec![0.0f32; seq_len * hidden];
let mut d_s = vec![0.0f32; value_heads * key_dim * value_dim];
let mut d_qkv_proj_all = vec![0.0f32; seq_len * qkv_dim];
for t in (0..seq_len).rev() {
let dy = &grad_outputs[t * hidden..(t + 1) * hidden];
let x_t = &saved.inputs[t * hidden..(t + 1) * hidden];
let mut d_gated_buf = vec![0.0f32; output_dim];
for j in 0..output_dim {
let mut acc = 0.0f64;
for i in 0..hidden {
acc += weights.out_proj[i * output_dim + j] as f64 * dy[i] as f64;
}
d_gated_buf[j] = acc as f32;
}
if have_lora {
let gated_buf_t = &saved.gated_buf[t * output_dim..(t + 1) * output_dim];
let h_out_t = &saved.h_out[t * lora_rank..(t + 1) * lora_rank];
let (gb, ga, dx_lora) = lora_vjp(
dy,
gated_buf_t,
h_out_t,
&saved.lora_a_out,
&saved.lora_b_out,
lora_rank,
output_dim,
hidden,
lora_scale,
);
for k in 0..ga.len() {
grad_a_out[k] += ga[k];
}
for k in 0..gb.len() {
grad_b_out[k] += gb[k];
}
for j in 0..output_dim {
d_gated_buf[j] += dx_lora[j];
}
}
let z_slice = &saved.z_proj[t * output_dim..(t + 1) * output_dim];
let mut d_o_heads = vec![0.0f32; output_dim];
let mut d_z_proj = vec![0.0f32; output_dim];
for h in 0..value_heads {
let rms = saved.rms_vals[t * value_heads + h];
let inv_rms = 1.0 / rms;
let o_off = (t * value_heads + h) * value_dim;
let o_h = &saved.o_heads[o_off..o_off + value_dim];
let sz_off = (t * value_heads + h) * value_dim;
let silu_z_h = &saved.silu_z[sz_off..sz_off + value_dim];
let d_g = &d_gated_buf[h * value_dim..(h + 1) * value_dim];
let d_o = &mut d_o_heads[h * value_dim..(h + 1) * value_dim];
let d_z = &mut d_z_proj[h * value_dim..(h + 1) * value_dim];
let mut d_xnorm = vec![0.0f32; value_dim];
for j in 0..value_dim {
let x_norm_j = o_h[j] * inv_rms;
d_xnorm[j] = d_g[j] * gamma[j] * silu_z_h[j];
let d_silu_z_j = d_g[j] * gamma[j] * x_norm_j;
let z_j = z_slice[h * value_dim + j];
let sig_z = sigmoid(z_j);
let d_silu_dz = sig_z * (1.0 + z_j * (1.0 - sig_z));
d_z[j] = d_silu_z_j * d_silu_dz;
}
let dot_dxnorm_o: f32 = d_xnorm.iter().zip(o_h.iter()).map(|(a, b)| a * b).sum();
let rms3_dim = rms * rms * rms * value_dim as f32;
for j in 0..value_dim {
d_o[j] = d_xnorm[j] * inv_rms - o_h[j] * dot_dxnorm_o / rms3_dim;
}
}
let dx_t = &mut grad_inputs[t * hidden..(t + 1) * hidden];
for j in 0..output_dim {
let dz_j = d_z_proj[j];
if dz_j.abs() < f32::EPSILON {
continue;
}
for i in 0..hidden {
dx_t[i] += weights.in_proj_z[j * hidden + i] * dz_j;
}
}
if have_lora {
let h_z_t = &saved.h_z[t * lora_rank..(t + 1) * lora_rank];
let (gb, ga, dx_lora) = lora_vjp(
&d_z_proj,
x_t,
h_z_t,
&saved.lora_a_z,
&saved.lora_b_z,
lora_rank,
hidden,
output_dim,
lora_scale,
);
for k in 0..ga.len() {
grad_a_z[k] += ga[k];
}
for k in 0..gb.len() {
grad_b_z[k] += gb[k];
}
for j in 0..hidden {
dx_t[j] += dx_lora[j];
}
}
let mut d_conv_out = vec![0.0f32; qkv_dim];
let mut d_alpha_vh = vec![0.0f32; value_heads]; let mut d_beta_raw_vh = vec![0.0f32; value_heads];
for h in 0..value_heads {
let kh = h / ratio;
let q_start = kh * key_dim;
let k_start = q_total + kh * key_dim;
let v_start = q_total * 2 + h * value_dim;
let g_h = saved.g[t * value_heads + h];
let beta_h = saved.beta[t * value_heads + h];
let q_off = (t * value_heads + h) * key_dim;
let k_off = (t * value_heads + h) * key_dim;
let kvm_off = (t * value_heads + h) * value_dim;
let sa_off = (t * value_heads + h) * key_dim * value_dim;
let ds_off = h * key_dim * value_dim;
let q_hat_h = &saved.q_hat[q_off..q_off + key_dim];
let k_hat_h = &saved.k_hat[k_off..k_off + key_dim];
let kv_mem_h = &saved.kv_mem[kvm_off..kvm_off + value_dim];
let s_t = &saved.s_after[sa_off..sa_off + key_dim * value_dim];
let d_o_h = &d_o_heads[h * value_dim..(h + 1) * value_dim];
let d_s_h = &mut d_s[ds_off..ds_off + key_dim * value_dim];
let mut d_q = vec![0.0f32; key_dim];
for i in 0..key_dim {
let mut acc = 0.0f32;
for j in 0..value_dim {
d_s_h[i * value_dim + j] += scale * d_o_h[j] * q_hat_h[i];
acc += s_t[i * value_dim + j] * d_o_h[j];
}
d_q[i] = scale * acc;
}
let v_off_h = (t * value_heads + h) * value_dim;
let v_h = &saved.v[v_off_h..v_off_h + value_dim];
let mut delta = vec![0.0f32; value_dim];
for j in 0..value_dim {
delta[j] = (v_h[j] - kv_mem_h[j] * g_h) * beta_h;
}
let s_prev_slice: Vec<f32> = if t == 0 {
vec![0.0f32; key_dim * value_dim]
} else {
let prev_off = ((t - 1) * value_heads + h) * key_dim * value_dim;
saved.s_after[prev_off..prev_off + key_dim * value_dim].to_vec()
};
let mut d_g_h: f32 = 0.0;
let mut d_k = vec![0.0f32; key_dim];
let mut d_delta = vec![0.0f32; value_dim];
for i in 0..key_dim {
for j in 0..value_dim {
let ds_ij = d_s_h[i * value_dim + j];
d_g_h += s_prev_slice[i * value_dim + j] * ds_ij;
d_k[i] += ds_ij * delta[j];
d_delta[j] += ds_ij * k_hat_h[i];
}
}
for ij in 0..key_dim * value_dim {
d_s_h[ij] *= g_h;
}
let mut d_v = vec![0.0f32; value_dim];
let mut d_kv_mem = vec![0.0f32; value_dim];
let mut d_beta_h: f32 = 0.0;
for j in 0..value_dim {
d_v[j] = d_delta[j] * beta_h;
d_kv_mem[j] = -d_delta[j] * g_h * beta_h;
d_g_h += (-kv_mem_h[j] * beta_h) * d_delta[j];
d_beta_h += (v_h[j] - kv_mem_h[j] * g_h) * d_delta[j];
}
for i in 0..key_dim {
for j in 0..value_dim {
d_s_h[i * value_dim + j] += d_kv_mem[j] * k_hat_h[i];
d_k[i] += s_prev_slice[i * value_dim + j] * d_kv_mem[j];
}
}
let alpha_h = saved.alpha_proj[t * value_heads + h];
let a_val = weights.a_log[h].exp();
let sp_arg = alpha_h + weights.dt_bias[h];
let sp_deriv = sigmoid(sp_arg); let d_alpha_from_g = d_g_h * g_h * (-a_val) * sp_deriv;
let q_eps_norm = saved.q_eps_norm[t * value_heads + h];
let k_eps_norm = saved.k_eps_norm[t * value_heads + h];
let dot_q = dot(&d_q, q_hat_h);
let dot_k = dot(&d_k, k_hat_h);
let mut d_q_raw = vec![0.0f32; key_dim];
let mut d_k_raw = vec![0.0f32; key_dim];
for i in 0..key_dim {
d_q_raw[i] = (d_q[i] - q_hat_h[i] * dot_q) / q_eps_norm;
d_k_raw[i] = (d_k[i] - k_hat_h[i] * dot_k) / k_eps_norm;
}
for i in 0..key_dim {
d_conv_out[q_start + i] += d_q_raw[i];
d_conv_out[k_start + i] += d_k_raw[i];
}
for j in 0..value_dim {
d_conv_out[v_start + j] += d_v[j];
}
let beta_raw_h = saved.beta_raw[t * value_heads + h];
let sig = sigmoid(beta_raw_h);
let d_beta_raw_h = d_beta_h * sig * (1.0 - sig);
d_alpha_vh[h] += d_alpha_from_g;
d_beta_raw_vh[h] += d_beta_raw_h;
}
let dx_t = &mut grad_inputs[t * hidden..(t + 1) * hidden];
for vh in 0..value_heads {
let d_a = d_alpha_vh[vh];
let d_b = d_beta_raw_vh[vh];
for i in 0..hidden {
dx_t[i] += weights.in_proj_a[vh * hidden + i] * d_a;
dx_t[i] += weights.in_proj_b[vh * hidden + i] * d_b;
}
}
if have_lora {
let h_b_t = &saved.h_b[t * lora_rank..(t + 1) * lora_rank];
let (gb, ga, dx_lora) = lora_vjp(
&d_beta_raw_vh,
x_t,
h_b_t,
&saved.lora_a_b,
&saved.lora_b_b,
lora_rank,
hidden,
value_heads,
lora_scale,
);
for k in 0..ga.len() {
grad_a_b[k] += ga[k];
}
for k in 0..gb.len() {
grad_b_b[k] += gb[k];
}
for j in 0..hidden {
dx_t[j] += dx_lora[j];
}
let h_a_t = &saved.h_a[t * lora_rank..(t + 1) * lora_rank];
let (gb, ga, dx_lora) = lora_vjp(
&d_alpha_vh,
x_t,
h_a_t,
&saved.lora_a_a,
&saved.lora_b_a,
lora_rank,
hidden,
value_heads,
lora_scale,
);
for k in 0..ga.len() {
grad_a_a[k] += ga[k];
}
for k in 0..gb.len() {
grad_b_a[k] += gb[k];
}
for j in 0..hidden {
dx_t[j] += dx_lora[j];
}
}
let cb_off = t * qkv_dim * buf_len;
let conv_buf_t = &saved.conv_buffers[cb_off..cb_off + qkv_dim * buf_len];
let qkv_in_t = &saved.qkv_proj[t * qkv_dim..(t + 1) * qkv_dim];
for ch in 0..qkv_dim {
let w_off = ch * kernel_size;
let buf_off = ch * buf_len;
let mut sum_pre = 0.0f32;
for tb in 0..buf_len {
sum_pre += conv_buf_t[buf_off + tb] * weights.conv1d_weight[w_off + tb];
}
sum_pre += qkv_in_t[ch] * weights.conv1d_weight[w_off + buf_len];
let sig = sigmoid(sum_pre);
let silu_deriv = sig * (1.0 + sum_pre * (1.0 - sig));
let d_sum = d_conv_out[ch] * silu_deriv;
d_qkv_proj_all[t * qkv_dim + ch] += d_sum * weights.conv1d_weight[w_off + buf_len];
for tb in 0..buf_len {
let lag = buf_len - tb; if t >= lag {
let src_t = t - lag;
d_qkv_proj_all[src_t * qkv_dim + ch] +=
d_sum * weights.conv1d_weight[w_off + tb];
}
}
}
let dx_t = &mut grad_inputs[t * hidden..(t + 1) * hidden];
for j in 0..qkv_dim {
let dq = d_qkv_proj_all[t * qkv_dim + j];
if dq.abs() < f32::EPSILON {
continue;
}
for i in 0..hidden {
dx_t[i] += weights.in_proj_qkv[j * hidden + i] * dq;
}
}
if have_lora {
let d_qkv_g = &d_qkv_proj_all[t * qkv_dim..(t + 1) * qkv_dim];
let h_qkv_t = &saved.h_qkv[t * lora_rank..(t + 1) * lora_rank];
let (gb, ga, dx_lora) = lora_vjp(
d_qkv_g,
x_t,
h_qkv_t,
&saved.lora_a_qkv,
&saved.lora_b_qkv,
lora_rank,
hidden,
qkv_dim,
lora_scale,
);
for k in 0..ga.len() {
grad_a_qkv[k] += ga[k];
}
for k in 0..gb.len() {
grad_b_qkv[k] += gb[k];
}
for j in 0..hidden {
dx_t[j] += dx_lora[j];
}
}
}
GdnGrads {
grad_a_qkv,
grad_b_qkv,
grad_a_z,
grad_b_z,
grad_a_b,
grad_b_b,
grad_a_a,
grad_b_a,
grad_a_out,
grad_b_out,
dx: grad_inputs,
}
}
#[inline]
fn silu_f32(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
#[inline]
fn l2_norm_sq(x: &[f32]) -> f32 {
x.iter().map(|v| v * v).sum()
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
fn conv1d_silu_fwd(
new_input: &[f32],
conv_buffer: &mut [f32],
conv_weight: &[f32],
output: &mut [f32],
conv_dim: usize,
kernel_size: usize,
) {
let buf_len = kernel_size.saturating_sub(1);
for ch in 0..conv_dim {
let w_off = ch * kernel_size;
let buf_off = ch * buf_len;
let row = &mut conv_buffer[buf_off..buf_off + buf_len];
let mut sum = 0.0f32;
for t in 0..buf_len {
sum += row[t] * conv_weight[w_off + t];
}
let x = new_input[ch];
sum += x * conv_weight[w_off + buf_len];
output[ch] = silu_f32(sum);
if buf_len > 1 {
row.copy_within(1..buf_len, 0);
}
if buf_len > 0 {
row[buf_len - 1] = x;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::attention::gdn::GatedDeltaNetWeights;
use crate::model::qwen35_config::Qwen35Config;
struct Rng(u64);
impl Rng {
fn new(seed: u64) -> Self {
Self(if seed == 0 {
0xDEAD_BEEF_CAFE_1234
} else {
seed
})
}
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn f32_range(&mut self, lo: f32, hi: f32) -> f32 {
let bits = (self.next() >> 40) as u32;
let t = (bits as f32) / ((1u32 << 24) as f32);
lo + t * (hi - lo)
}
fn fill(&mut self, out: &mut [f32], lo: f32, hi: f32) {
for v in out {
*v = self.f32_range(lo, hi);
}
}
}
fn make_tiny_weights(
hidden: usize,
num_kh: usize,
num_vh: usize,
key_dim: usize,
value_dim: usize,
kernel_size: usize,
seed: u64,
) -> GatedDeltaNetWeights {
let qkv_dim = num_kh * key_dim * 2 + num_vh * value_dim; let output_dim = num_vh * value_dim;
let mut rng = Rng::new(seed);
let scale = 0.05;
let mut in_proj_qkv = vec![0.0f32; qkv_dim * hidden];
let mut in_proj_z = vec![0.0f32; output_dim * hidden];
let mut in_proj_b = vec![0.0f32; num_vh * hidden];
let mut in_proj_a = vec![0.0f32; num_vh * hidden];
let mut conv1d_weight = vec![0.0f32; qkv_dim * kernel_size];
let mut out_proj = vec![0.0f32; hidden * output_dim];
let mut norm_weight = vec![0.0f32; value_dim];
let mut a_log = vec![0.0f32; num_vh];
let mut dt_bias = vec![0.0f32; num_vh];
rng.fill(&mut in_proj_qkv, -scale, scale);
rng.fill(&mut in_proj_z, -scale, scale);
rng.fill(&mut in_proj_b, -scale, scale);
rng.fill(&mut in_proj_a, -scale, scale);
rng.fill(&mut conv1d_weight, -scale, scale);
rng.fill(&mut out_proj, -scale, scale);
for g in &mut norm_weight {
*g = rng.f32_range(0.9, 1.1);
}
for a in &mut a_log {
*a = rng.f32_range(-0.3, 2.0);
}
for dt in &mut dt_bias {
*dt = rng.f32_range(-0.1, 0.1);
}
GatedDeltaNetWeights {
in_proj_qkv,
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z,
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b,
in_proj_b_rows: num_vh,
in_proj_b_cols: hidden,
in_proj_a,
in_proj_a_rows: num_vh,
in_proj_a_cols: hidden,
a_log,
dt_bias,
conv1d_weight,
conv_dim: qkv_dim,
kernel_size,
norm_weight,
out_proj,
out_proj_rows: hidden,
out_proj_cols: output_dim,
}
}
fn tiny_cfg(
hidden: usize,
num_kh: usize,
num_vh: usize,
key_dim: usize,
value_dim: usize,
kernel_size: usize,
) -> Qwen35Config {
let mut cfg = Qwen35Config::qwen35_2b();
cfg.hidden_size = hidden;
cfg.linear_num_key_heads = num_kh;
cfg.linear_num_value_heads = Some(num_vh);
cfg.linear_key_head_dim = key_dim;
cfg.linear_value_head_dim = value_dim;
cfg.linear_conv_kernel_dim = kernel_size;
cfg
}
fn linear_loss(outputs: &[f32], coeffs: &[f32]) -> f64 {
outputs
.iter()
.zip(coeffs.iter())
.map(|(&y, &w)| (y as f64) * (w as f64))
.sum()
}
fn fd_grad_inputs_linear(
inputs: &[f32],
weights: &GatedDeltaNetWeights,
cfg: &Qwen35Config,
seq_len: usize,
hidden: usize,
eps: f64,
coeffs: &[f32],
) -> Vec<f32> {
let num_kh = cfg.linear_num_key_heads;
let value_heads = cfg.linear_num_value_heads();
let key_dim = cfg.linear_key_head_dim;
let value_dim = cfg.linear_value_head_dim;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let kernel_size = cfg.linear_conv_kernel_dim;
let scale = 1.0 / (key_dim as f32).sqrt();
let mut grad = vec![0.0f32; inputs.len()];
for idx in 0..inputs.len() {
let mut inputs_p: Vec<f64> = inputs.iter().map(|&v| v as f64).collect();
let mut inputs_m: Vec<f64> = inputs.iter().map(|&v| v as f64).collect();
inputs_p[idx] += eps;
inputs_m[idx] -= eps;
let inp_p_f32: Vec<f32> = inputs_p.iter().map(|&v| v as f32).collect();
let inp_m_f32: Vec<f32> = inputs_m.iter().map(|&v| v as f32).collect();
let mut saved_p = GdnSaved::new(
seq_len,
num_kh,
value_heads,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut saved_m = GdnSaved::new(
seq_len,
num_kh,
value_heads,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut out_p = vec![0.0f32; seq_len * hidden];
let mut out_m = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
&inp_p_f32,
weights,
cfg,
&mut saved_p,
&mut out_p,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
0,
0.0,
);
gdn_forward_save(
&inp_m_f32,
weights,
cfg,
&mut saved_m,
&mut out_m,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
0,
0.0,
);
let lp = linear_loss(&out_p, coeffs);
let lm = linear_loss(&out_m, coeffs);
grad[idx] = ((lp - lm) / (2.0 * eps)) as f32;
}
grad
}
fn run_gradcheck_linear(
hidden: usize,
num_kh: usize,
num_vh: usize,
key_dim: usize,
value_dim: usize,
kernel_size: usize,
seq_len: usize,
input_scale: f32,
weight_seed: u64,
input_seed: u64,
coeff_seed: u64,
eps: f64,
) -> (f32, usize, f32, f32) {
let cfg = tiny_cfg(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size);
let weights = make_tiny_weights(
hidden,
num_kh,
num_vh,
key_dim,
value_dim,
kernel_size,
weight_seed,
);
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let scale = 1.0 / (key_dim as f32).sqrt();
let mut rng_in = Rng::new(input_seed);
let mut inputs = vec![0.0f32; seq_len * hidden];
rng_in.fill(&mut inputs, -input_scale, input_scale);
let mut rng_c = Rng::new(coeff_seed);
let mut coeffs = vec![0.0f32; seq_len * hidden];
rng_c.fill(&mut coeffs, -1.0, 1.0);
let mut saved = GdnSaved::new(
seq_len,
num_kh,
num_vh,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut outputs = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
&inputs,
&weights,
&cfg,
&mut saved,
&mut outputs,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
0,
0.0,
);
let gdn_grads = gdn_backward(&coeffs, &saved, &weights);
let analytic = gdn_grads.dx;
let fd = fd_grad_inputs_linear(&inputs, &weights, &cfg, seq_len, hidden, eps, &coeffs);
let mut max_rel = 0.0f32;
let mut worst = 0;
let mut n_tested = 0usize;
for i in 0..analytic.len() {
let abs_err = (analytic[i] - fd[i]).abs();
let mag = fd[i].abs().max(analytic[i].abs());
if mag < 1e-4 {
continue;
}
n_tested += 1;
let rel = abs_err / mag;
if rel > max_rel {
max_rel = rel;
worst = i;
}
}
assert!(
n_tested >= 10,
"gradcheck skipped too many entries (only {n_tested} with |g| >= 1e-4); \
increase input_scale or weight_scale to produce meaningful gradients"
);
(max_rel, worst, analytic[worst], fd[worst])
}
#[test]
fn gradcheck_gdn_backward() {
let (max_rel, worst, analytic_v, fd_v) = run_gradcheck_linear(
32, 1, 1, 8, 8, 3, 8, 0.5, 42, 1337, 999, 1e-3, );
assert!(
max_rel < 1e-2,
"gradcheck FAILED: max rel-err = {max_rel:.2e} at index {worst} \
(analytic={analytic_v}, fd={fd_v})",
);
}
#[test]
fn gradcheck_gdn_backward_multi_head() {
let (max_rel, worst, analytic_v, fd_v) = run_gradcheck_linear(
32, 2, 4, 6, 6, 2, 4, 0.5, 99, 2024, 777, 1e-3, );
assert!(
max_rel < 1e-2,
"multi-head gradcheck FAILED: max rel-err = {max_rel:.2e} at index {worst} \
(analytic={analytic_v}, fd={fd_v})",
);
}
#[test]
fn zero_grad_output_yields_zero_input_grad() {
let hidden = 16;
let num_kh = 1;
let num_vh = 1;
let key_dim = 4;
let value_dim = 4;
let kernel_size = 2;
let seq_len = 3;
let cfg = tiny_cfg(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size);
let weights = make_tiny_weights(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size, 7);
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let scale = 1.0 / (key_dim as f32).sqrt();
let mut rng = Rng::new(888);
let mut inputs = vec![0.0f32; seq_len * hidden];
rng.fill(&mut inputs, -0.5, 0.5);
let mut saved = GdnSaved::new(
seq_len,
num_kh,
num_vh,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut outputs = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
&inputs,
&weights,
&cfg,
&mut saved,
&mut outputs,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
0,
0.0,
);
let grad_out = vec![0.0f32; seq_len * hidden];
let gdn_grads = gdn_backward(&grad_out, &saved, &weights);
let analytic = gdn_grads.dx;
for (i, &v) in analytic.iter().enumerate() {
assert!(
v.abs() < 1e-10,
"zero upstream grad should give zero input grad at [{i}], got {v}"
);
}
}
#[allow(clippy::too_many_arguments)]
fn run_lora_forward_backward(
inputs: &[f32],
weights: &GatedDeltaNetWeights,
cfg: &Qwen35Config,
seq_len: usize,
coeffs: &[f32],
lora_rank: usize,
lora_scale: f32,
a_qkv: &[f32],
b_qkv: &[f32],
a_z: &[f32],
b_z: &[f32],
a_b: &[f32],
b_b: &[f32],
a_a: &[f32],
b_a: &[f32],
a_out: &[f32],
b_out: &[f32],
) -> GdnGrads {
let hidden = cfg.hidden_size;
let num_kh = cfg.linear_num_key_heads;
let num_vh = cfg.linear_num_value_heads();
let key_dim = cfg.linear_key_head_dim;
let value_dim = cfg.linear_value_head_dim;
let kernel_size = cfg.linear_conv_kernel_dim;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let scale = 1.0 / (key_dim as f32).sqrt();
let mut saved = GdnSaved::new(
seq_len,
num_kh,
num_vh,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut outputs = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
inputs,
weights,
cfg,
&mut saved,
&mut outputs,
Some(a_qkv),
Some(b_qkv),
Some(a_z),
Some(b_z),
Some(a_b),
Some(b_b),
Some(a_a),
Some(b_a),
Some(a_out),
Some(b_out),
lora_rank,
lora_scale,
);
gdn_backward(coeffs, &saved, weights)
}
#[allow(clippy::too_many_arguments)]
fn fd_lora_weight_entry(
inputs: &[f32],
weights: &GatedDeltaNetWeights,
cfg: &Qwen35Config,
seq_len: usize,
coeffs: &[f32],
lora_rank: usize,
lora_scale: f32,
a_qkv: &[f32],
b_qkv: &[f32],
a_z: &[f32],
b_z: &[f32],
a_b: &[f32],
b_b: &[f32],
a_a: &[f32],
b_a: &[f32],
a_out: &[f32],
b_out: &[f32],
which: usize, idx: usize,
eps: f32,
) -> f32 {
let perturbed_loss = |delta: f32| -> f64 {
let mut aq = a_qkv.to_vec();
let mut bq = b_qkv.to_vec();
let mut az = a_z.to_vec();
let mut bz = b_z.to_vec();
let mut ab = a_b.to_vec();
let mut bb = b_b.to_vec();
let mut aa = a_a.to_vec();
let mut ba = b_a.to_vec();
let mut ao = a_out.to_vec();
let mut bo = b_out.to_vec();
let arr: &mut Vec<f32> = match which {
0 => &mut aq,
1 => &mut bq,
2 => &mut az,
3 => &mut bz,
4 => &mut ab,
5 => &mut bb,
6 => &mut aa,
7 => &mut ba,
8 => &mut ao,
_ => &mut bo,
};
arr[idx] += delta;
let hidden = cfg.hidden_size;
let num_kh = cfg.linear_num_key_heads;
let num_vh = cfg.linear_num_value_heads();
let key_dim = cfg.linear_key_head_dim;
let value_dim = cfg.linear_value_head_dim;
let kernel_size = cfg.linear_conv_kernel_dim;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let scale = 1.0 / (key_dim as f32).sqrt();
let mut saved = GdnSaved::new(
seq_len,
num_kh,
num_vh,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut outputs = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
inputs,
weights,
cfg,
&mut saved,
&mut outputs,
Some(&aq),
Some(&bq),
Some(&az),
Some(&bz),
Some(&ab),
Some(&bb),
Some(&aa),
Some(&ba),
Some(&ao),
Some(&bo),
lora_rank,
lora_scale,
);
outputs
.iter()
.zip(coeffs.iter())
.map(|(&y, &c)| (y as f64) * (c as f64))
.sum()
};
let lp = perturbed_loss(eps);
let lm = perturbed_loss(-eps);
((lp - lm) / (2.0 * eps as f64)) as f32
}
const LORA_ARRAY_NAMES: [&str; 10] = [
"grad_a_qkv",
"grad_b_qkv",
"grad_a_z",
"grad_b_z",
"grad_a_b",
"grad_b_b",
"grad_a_a",
"grad_b_a",
"grad_a_out",
"grad_b_out",
];
struct PerArrayCoverage {
per_array: [(usize, f32); 10],
overall_max_rel: f32,
overall_n_tested: usize,
}
#[allow(clippy::too_many_arguments)]
fn run_lora_weight_gradcheck(
hidden: usize,
num_kh: usize,
num_vh: usize,
key_dim: usize,
value_dim: usize,
kernel_size: usize,
seq_len: usize,
lora_rank: usize,
lora_scale: f32,
weight_seed: u64,
input_seed: u64,
lora_seed: u64,
coeff_seed: u64,
eps: f32,
) -> PerArrayCoverage {
let cfg = tiny_cfg(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size);
let weights = make_tiny_weights(
hidden,
num_kh,
num_vh,
key_dim,
value_dim,
kernel_size,
weight_seed,
);
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let mut rng_in = Rng::new(input_seed);
let mut inputs = vec![0.0f32; seq_len * hidden];
rng_in.fill(&mut inputs, -0.5, 0.5);
let mut rng_c = Rng::new(coeff_seed);
let mut coeffs = vec![0.0f32; seq_len * hidden];
rng_c.fill(&mut coeffs, -1.0, 1.0);
let mut rng_l = Rng::new(lora_seed);
let mut a_qkv = vec![0.0f32; lora_rank * hidden]; let mut b_qkv = vec![0.0f32; qkv_dim * lora_rank]; let mut a_z = vec![0.0f32; lora_rank * hidden]; let mut b_z = vec![0.0f32; output_dim * lora_rank]; let mut a_b = vec![0.0f32; lora_rank * hidden]; let mut b_b = vec![0.0f32; num_vh * lora_rank]; let mut a_a = vec![0.0f32; lora_rank * hidden]; let mut b_a = vec![0.0f32; num_vh * lora_rank]; let mut a_out = vec![0.0f32; lora_rank * output_dim]; let mut b_out = vec![0.0f32; hidden * lora_rank]; for arr in [
&mut a_qkv, &mut b_qkv, &mut a_z, &mut b_z, &mut a_b, &mut b_b, &mut a_a, &mut b_a,
&mut a_out, &mut b_out,
]
.iter_mut()
{
rng_l.fill(arr, -0.05, 0.05);
}
let grads = run_lora_forward_backward(
&inputs, &weights, &cfg, seq_len, &coeffs, lora_rank, lora_scale, &a_qkv, &b_qkv, &a_z,
&b_z, &a_b, &b_b, &a_a, &b_a, &a_out, &b_out,
);
let analytic_arrays: [(&[f32], usize); 10] = [
(&grads.grad_a_qkv, a_qkv.len()),
(&grads.grad_b_qkv, b_qkv.len()),
(&grads.grad_a_z, a_z.len()),
(&grads.grad_b_z, b_z.len()),
(&grads.grad_a_b, a_b.len()),
(&grads.grad_b_b, b_b.len()),
(&grads.grad_a_a, a_a.len()),
(&grads.grad_b_a, b_a.len()),
(&grads.grad_a_out, a_out.len()),
(&grads.grad_b_out, b_out.len()),
];
let mut overall_max_rel = 0.0f32;
let mut overall_n_tested = 0usize;
let mut per_array = [(0usize, 0.0f32); 10];
for (which, (analytic_slice, arr_len)) in analytic_arrays.iter().enumerate() {
let step = (arr_len / 6).max(1);
let mut idx = 0usize;
let mut arr_n_tested = 0usize;
let mut arr_max_rel = 0.0f32;
while idx < *arr_len {
let analytic_v = analytic_slice[idx];
let fd_v = fd_lora_weight_entry(
&inputs, &weights, &cfg, seq_len, &coeffs, lora_rank, lora_scale, &a_qkv,
&b_qkv, &a_z, &b_z, &a_b, &b_b, &a_a, &b_a, &a_out, &b_out, which, idx, eps,
);
let mag = fd_v.abs().max(analytic_v.abs());
if mag >= 1e-4 {
arr_n_tested += 1;
overall_n_tested += 1;
let rel = (fd_v - analytic_v).abs() / mag;
if rel > arr_max_rel {
arr_max_rel = rel;
}
if rel > overall_max_rel {
overall_max_rel = rel;
}
}
idx += step;
}
per_array[which] = (arr_n_tested, arr_max_rel);
}
PerArrayCoverage {
per_array,
overall_max_rel,
overall_n_tested,
}
}
#[test]
fn gradcheck_gdn_lora_weight_grads() {
let cov = run_lora_weight_gradcheck(
32, 1, 1, 8, 8, 3, 14, 2, 16.0, 42, 1337, 777, 999, 1e-3, );
for (name, (n, rel)) in LORA_ARRAY_NAMES.iter().zip(cov.per_array.iter()) {
println!(" {name:<12}: n_tested={n} max_rel={rel:.3e}");
}
for (name, (n, _)) in LORA_ARRAY_NAMES.iter().zip(cov.per_array.iter()) {
assert!(
*n >= 1,
"lora weight gradcheck: array {name} has ZERO testable entries \
above the 1e-4 floor — this array is not actually being verified"
);
}
assert!(
cov.overall_n_tested >= 10,
"lora weight gradcheck: too few testable entries ({}); \
increase array sizes or tighten threshold",
cov.overall_n_tested
);
assert!(
cov.overall_max_rel < 1e-2,
"gdn lora weight gradcheck failed: max_rel={:.3e} (n_tested={})",
cov.overall_max_rel,
cov.overall_n_tested
);
}
#[test]
fn gradcheck_gdn_lora_weight_grads_multi_head() {
let cov = run_lora_weight_gradcheck(
32, 2, 4, 8, 8, 3,
20, 2, 12.0, 99, 2024, 555, 111, 1e-3, );
for (name, (n, rel)) in LORA_ARRAY_NAMES.iter().zip(cov.per_array.iter()) {
println!(" {name:<12}: n_tested={n} max_rel={rel:.3e}");
}
for (name, (n, _)) in LORA_ARRAY_NAMES.iter().zip(cov.per_array.iter()) {
assert!(
*n >= 1,
"lora weight gradcheck (multi-head): array {name} has ZERO testable \
entries above the 1e-4 floor — this array is not actually being verified"
);
}
assert!(
cov.overall_n_tested >= 10,
"lora weight gradcheck (multi-head): too few testable entries ({})",
cov.overall_n_tested
);
assert!(
cov.overall_max_rel < 1e-2,
"gdn lora weight gradcheck (multi-head) failed: max_rel={:.3e} (n_tested={})",
cov.overall_max_rel,
cov.overall_n_tested
);
}
#[test]
fn source_parity_forward_save_vs_fused_asymmetric_heads() {
use crate::attention::gdn::GatedDeltaNetState;
use crate::attention::gdn_fused::{GatedDeltaNetFusedScratch, gated_delta_net_step_fused};
use crate::lora_hook::NoopLoraHook;
let hidden = 32;
let num_kh = 2;
let num_vh = 4; let key_dim = 8;
let value_dim = 8;
let kernel_size = 3;
let seq_len = 6;
let cfg = tiny_cfg(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size);
let weights =
make_tiny_weights(hidden, num_kh, num_vh, key_dim, value_dim, kernel_size, 321);
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let scale = 1.0 / (key_dim as f32).sqrt();
let mut rng = Rng::new(654);
let mut inputs = vec![0.0f32; seq_len * hidden];
rng.fill(&mut inputs, -0.5, 0.5);
let mut saved = GdnSaved::new(
seq_len,
num_kh,
num_vh,
key_dim,
value_dim,
hidden,
qkv_dim,
output_dim,
kernel_size,
scale,
cfg.rms_norm_eps,
);
let mut outputs_a = vec![0.0f32; seq_len * hidden];
gdn_forward_save(
&inputs,
&weights,
&cfg,
&mut saved,
&mut outputs_a,
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
0,
0.0,
);
let mut state = GatedDeltaNetState::new(&cfg);
let mut scratch = GatedDeltaNetFusedScratch::default();
let hook = NoopLoraHook;
let mut outputs_b = vec![0.0f32; seq_len * hidden];
for t in 0..seq_len {
let x_t = &inputs[t * hidden..(t + 1) * hidden];
let y_t = &mut outputs_b[t * hidden..(t + 1) * hidden];
gated_delta_net_step_fused(
x_t,
&mut state,
&weights,
&cfg,
&mut scratch,
y_t,
&hook,
0,
);
}
let mut max_abs_diff = 0.0f32;
let mut worst = 0;
for i in 0..outputs_a.len() {
let d = (outputs_a[i] - outputs_b[i]).abs();
if d > max_abs_diff {
max_abs_diff = d;
worst = i;
}
}
assert!(
max_abs_diff < 1e-4,
"gdn_forward_save diverges from the shipping gated_delta_net_step_fused \
on an asymmetric-head fixture (num_kh={num_kh}, num_vh={num_vh}): \
max_abs_diff={max_abs_diff:.3e} at flat index {worst} \
(a={}, b={})",
outputs_a[worst],
outputs_b[worst],
);
}
}