use ferrox_core::matmul::{rms_norm, silu};
use ferrox_core::weight_matrix::WeightMatrix;
#[derive(Debug, Clone, Copy)]
pub struct GdnConfig {
pub hidden_dim: usize,
pub num_key_heads: usize,
pub num_value_heads: usize,
pub key_head_dim: usize,
pub value_head_dim: usize,
pub conv_kernel_size: usize,
pub rms_norm_eps: f32,
}
impl GdnConfig {
pub fn key_dim(&self) -> usize {
self.num_key_heads * self.key_head_dim
}
pub fn value_dim(&self) -> usize {
self.num_value_heads * self.value_head_dim
}
pub fn qkv_dim(&self) -> usize {
2 * self.key_dim() + self.value_dim()
}
pub fn heads_per_key_group(&self) -> usize {
self.num_value_heads / self.num_key_heads
}
}
pub struct GdnWeights {
pub attn_qkv: WeightMatrix, pub attn_gate: WeightMatrix, pub ssm_conv1d: Vec<f32>,
pub ssm_dt: Vec<f32>, pub ssm_a: Vec<f32>, pub ssm_beta: WeightMatrix, pub ssm_alpha: WeightMatrix, pub ssm_norm: Vec<f32>, pub ssm_out: WeightMatrix, }
pub struct GdnState {
conv_hist: Vec<f32>,
recurrent: Vec<f32>,
}
impl GdnState {
pub fn new(cfg: &GdnConfig) -> Self {
Self {
conv_hist: Vec::new(),
recurrent: vec![0f32; cfg.num_value_heads * cfg.value_head_dim * cfg.key_head_dim],
}
}
}
fn softplus(x: f32) -> f32 {
if x > 20.0 {
x
} else {
(1.0 + x.exp()).ln()
}
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn l2_normalize(v: &mut [f32], eps: f32) {
let norm_sq: f32 = v.iter().map(|x| x * x).sum();
let scale = 1.0 / (norm_sq + eps).sqrt();
for x in v.iter_mut() {
*x *= scale;
}
}
fn causal_conv_step(
weight: &[f32],
history: &mut Vec<f32>,
current: &[f32],
kernel_size: usize,
dim: usize,
) -> Vec<f32> {
let hist_len = history.len() / dim.max(1);
let missing = (kernel_size - 1).saturating_sub(hist_len);
let mut y = vec![0f32; dim];
for j in 0..kernel_size {
if j < missing {
continue;
}
let src: &[f32] = if j == kernel_size - 1 {
current
} else {
let hist_idx = j - missing;
&history[hist_idx * dim..(hist_idx + 1) * dim]
};
for d in 0..dim {
y[d] += weight[d * kernel_size + j] * src[d];
}
}
for v in y.iter_mut() {
*v = silu(*v);
}
history.extend_from_slice(current);
let max_hist_len = (kernel_size - 1) * dim;
if history.len() > max_hist_len {
let excess = history.len() - max_hist_len;
history.drain(0..excess);
}
y
}
pub fn gdn_forward_token(
weights: &GdnWeights,
cfg: &GdnConfig,
hidden: &[f32],
state: &mut GdnState,
) -> Vec<f32> {
assert_eq!(hidden.len(), cfg.hidden_dim);
assert!(
cfg.num_key_heads > 0 && cfg.num_value_heads.is_multiple_of(cfg.num_key_heads),
"GDN needs num_value_heads ({}) to be a positive multiple of num_key_heads ({}); \
otherwise repeat_interleave has no whole replication factor and some V heads would \
silently read a K head that never fed them",
cfg.num_value_heads,
cfg.num_key_heads
);
let qkv_dim = cfg.qkv_dim();
let key_dim = cfg.key_dim();
let value_dim = cfg.value_dim();
let key_head_dim = cfg.key_head_dim;
let value_head_dim = cfg.value_head_dim;
let rep = cfg.heads_per_key_group();
let qkv_lin = weights.attn_qkv.apply(hidden);
let z = weights.attn_gate.apply(hidden);
let beta_raw = weights.ssm_beta.apply(hidden);
let alpha_raw = weights.ssm_alpha.apply(hidden);
let qkv = causal_conv_step(
&weights.ssm_conv1d,
&mut state.conv_hist,
&qkv_lin,
cfg.conv_kernel_size,
qkv_dim,
);
let (q_all, rest) = qkv.split_at(key_dim);
let (k_all, v_all) = rest.split_at(key_dim);
let scale = 1.0 / (key_head_dim as f32).sqrt();
let mut y_flat = vec![0f32; value_dim];
#[allow(clippy::needless_range_loop)]
for h in 0..cfg.num_value_heads {
let k_base = (h / rep) * key_head_dim;
let v_base = h * value_head_dim;
let mut q_h = q_all[k_base..k_base + key_head_dim].to_vec();
let mut k_h = k_all[k_base..k_base + key_head_dim].to_vec();
let v_h = &v_all[v_base..v_base + value_head_dim];
l2_normalize(&mut q_h, 1e-6);
l2_normalize(&mut k_h, 1e-6);
for x in q_h.iter_mut() {
*x *= scale;
}
let gate = softplus(alpha_raw[h] + weights.ssm_dt[h]) * weights.ssm_a[h];
let decay = gate.exp();
let beta = sigmoid(beta_raw[h]);
let block = value_head_dim * key_head_dim;
let s_base = h * block;
let s = &mut state.recurrent[s_base..s_base + block];
for cell in s.iter_mut() {
*cell *= decay;
}
let mut kv_mem = vec![0f32; value_head_dim];
for v_idx in 0..value_head_dim {
let mut acc = 0f32;
for k_idx in 0..key_head_dim {
acc += s[v_idx * key_head_dim + k_idx] * k_h[k_idx];
}
kv_mem[v_idx] = acc;
}
for v_idx in 0..value_head_dim {
let delta = (v_h[v_idx] - kv_mem[v_idx]) * beta;
for k_idx in 0..key_head_dim {
s[v_idx * key_head_dim + k_idx] += delta * k_h[k_idx];
}
}
for v_idx in 0..value_head_dim {
let mut acc = 0f32;
for k_idx in 0..key_head_dim {
acc += s[v_idx * key_head_dim + k_idx] * q_h[k_idx];
}
y_flat[v_base + v_idx] = acc;
}
}
let mut gated = vec![0f32; value_dim];
for h in 0..cfg.num_value_heads {
let base = h * value_head_dim;
let normed = rms_norm(
&y_flat[base..base + value_head_dim],
&weights.ssm_norm,
cfg.rms_norm_eps,
);
for i in 0..value_head_dim {
gated[base + i] = silu(z[base + i]) * normed[i];
}
}
weights.ssm_out.apply(&gated)
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_core::tensor::Tensor;
const HIDDEN: usize = 4;
const N_HEADS: usize = 2;
const HEAD_DIM: usize = 2;
const CONV_K: usize = 2;
const QKV_DIM: usize = 3 * N_HEADS * HEAD_DIM; const V_DIM: usize = N_HEADS * HEAD_DIM;
fn wm(data: &[f32], rows: usize, cols: usize) -> WeightMatrix {
assert_eq!(data.len(), rows * cols);
WeightMatrix::F32(Tensor::new(data.to_vec(), vec![rows, cols]))
}
fn cfg() -> GdnConfig {
GdnConfig {
hidden_dim: HIDDEN,
num_key_heads: N_HEADS,
num_value_heads: N_HEADS,
key_head_dim: HEAD_DIM,
value_head_dim: HEAD_DIM,
conv_kernel_size: CONV_K,
rms_norm_eps: 1e-5,
}
}
fn make_weights() -> GdnWeights {
let mut qkv = Vec::with_capacity(QKV_DIM * HIDDEN);
for i in 0..QKV_DIM * HIDDEN {
qkv.push(((i % 7) as f32 - 3.0) * 0.1);
}
let mut gate = Vec::with_capacity(V_DIM * HIDDEN);
for i in 0..V_DIM * HIDDEN {
gate.push(((i % 5) as f32 - 2.0) * 0.08);
}
let mut conv = Vec::with_capacity(QKV_DIM * CONV_K);
for i in 0..QKV_DIM * CONV_K {
conv.push(if i % CONV_K == CONV_K - 1 { 1.0 } else { 0.1 });
}
let mut beta = Vec::with_capacity(N_HEADS * HIDDEN);
let mut alpha = Vec::with_capacity(N_HEADS * HIDDEN);
for i in 0..N_HEADS * HIDDEN {
beta.push(((i % 3) as f32 - 1.0) * 0.2);
alpha.push(((i % 4) as f32 - 1.5) * 0.15);
}
let mut out = Vec::with_capacity(HIDDEN * V_DIM);
for i in 0..HIDDEN * V_DIM {
out.push(((i % 6) as f32 - 2.5) * 0.12);
}
GdnWeights {
attn_qkv: wm(&qkv, QKV_DIM, HIDDEN),
attn_gate: wm(&gate, V_DIM, HIDDEN),
ssm_conv1d: conv,
ssm_dt: vec![0.1, -0.05],
ssm_a: vec![-0.5, -0.75],
ssm_beta: wm(&beta, N_HEADS, HIDDEN),
ssm_alpha: wm(&alpha, N_HEADS, HIDDEN),
ssm_norm: vec![1.0, 1.0],
ssm_out: wm(&out, HIDDEN, V_DIM),
}
}
struct RawGdn {
qkv: Vec<f32>, gate: Vec<f32>, conv: Vec<f32>, dt: Vec<f32>, a: Vec<f32>, beta: Vec<f32>, alpha: Vec<f32>, norm: Vec<f32>, out: Vec<f32>, }
impl RawGdn {
fn to_weights(&self, cfg: &GdnConfig) -> GdnWeights {
let h = cfg.hidden_dim;
GdnWeights {
attn_qkv: wm(&self.qkv, cfg.qkv_dim(), h),
attn_gate: wm(&self.gate, cfg.value_dim(), h),
ssm_conv1d: self.conv.clone(),
ssm_dt: self.dt.clone(),
ssm_a: self.a.clone(),
ssm_beta: wm(&self.beta, cfg.num_value_heads, h),
ssm_alpha: wm(&self.alpha, cfg.num_value_heads, h),
ssm_norm: self.norm.clone(),
ssm_out: wm(&self.out, h, cfg.value_dim()),
}
}
}
fn fill(n: usize, seed: usize) -> Vec<f32> {
(0..n)
.map(|i| {
let k = (i * 37 + seed * 101) % 23;
(k as f32 - 11.0)
* 0.043
* if (i + seed).is_multiple_of(2) {
1.0
} else {
-1.0
}
})
.collect()
}
fn matvec(rows_data: &[f32], rows: usize, cols: usize, x: &[f32]) -> Vec<f32> {
assert_eq!(rows_data.len(), rows * cols);
assert_eq!(x.len(), cols);
(0..rows)
.map(|r| (0..cols).map(|c| rows_data[r * cols + c] * x[c]).sum())
.collect()
}
fn ref_sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn ref_softplus(x: f32) -> f32 {
(1.0 + x.exp()).ln()
}
fn ref_silu(x: f32) -> f32 {
x * ref_sigmoid(x)
}
fn ref_l2norm(v: &[f32]) -> Vec<f32> {
let sum_sq: f32 = v.iter().map(|x| x * x).sum();
let inv = 1.0 / (sum_sq + 1e-6).sqrt();
v.iter().map(|x| x * inv).collect()
}
fn reference_forward(raw: &RawGdn, cfg: &GdnConfig, tokens: &[Vec<f32>]) -> Vec<Vec<f32>> {
let hidden = cfg.hidden_dim;
let qkv_dim = cfg.qkv_dim();
let key_dim = cfg.key_dim();
let value_dim = cfg.value_dim();
let dk = cfg.key_head_dim;
let dv = cfg.value_head_dim;
let kernel = cfg.conv_kernel_size;
let rep = cfg.num_value_heads / cfg.num_key_heads;
let scale = 1.0 / (dk as f32).sqrt();
let mixed: Vec<Vec<f32>> = tokens
.iter()
.map(|h| matvec(&raw.qkv, qkv_dim, hidden, h))
.collect();
let mut conved = vec![vec![0f32; qkv_dim]; tokens.len()];
for (t, conved_t) in conved.iter_mut().enumerate() {
for (d, out_d) in conved_t.iter_mut().enumerate() {
let mut acc = 0f32;
for j in 0..kernel {
let src = t as isize - (kernel as isize - 1) + j as isize;
if src < 0 {
continue;
}
acc += raw.conv[d * kernel + j] * mixed[src as usize][d];
}
*out_d = ref_silu(acc);
}
}
let mut state = vec![vec![vec![0f32; dv]; dk]; cfg.num_value_heads];
let mut outputs = Vec::with_capacity(tokens.len());
for (t, token) in tokens.iter().enumerate() {
let z = matvec(&raw.gate, value_dim, hidden, token);
let a_raw = matvec(&raw.alpha, cfg.num_value_heads, hidden, token);
let b_raw = matvec(&raw.beta, cfg.num_value_heads, hidden, token);
let q_slice = &conved[t][0..key_dim];
let k_slice = &conved[t][key_dim..2 * key_dim];
let v_slice = &conved[t][2 * key_dim..];
let mut q_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
let mut k_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
for kh in 0..cfg.num_key_heads {
for _ in 0..rep {
q_heads.push(q_slice[kh * dk..(kh + 1) * dk].to_vec());
k_heads.push(k_slice[kh * dk..(kh + 1) * dk].to_vec());
}
}
let mut core = vec![0f32; value_dim];
for h in 0..cfg.num_value_heads {
let q = ref_l2norm(&q_heads[h]);
let k = ref_l2norm(&k_heads[h]);
let v = &v_slice[h * dv..(h + 1) * dv];
let decay = (ref_softplus(a_raw[h] + raw.dt[h]) * raw.a[h]).exp();
let beta = ref_sigmoid(b_raw[h]);
for row in state[h].iter_mut() {
for cell in row.iter_mut() {
*cell *= decay;
}
}
let mut kv_mem = vec![0f32; dv];
for (k_idx, row) in state[h].iter().enumerate() {
for (v_idx, cell) in row.iter().enumerate() {
kv_mem[v_idx] += cell * k[k_idx];
}
}
let delta: Vec<f32> = (0..dv).map(|i| (v[i] - kv_mem[i]) * beta).collect();
for (k_idx, row) in state[h].iter_mut().enumerate() {
for (v_idx, cell) in row.iter_mut().enumerate() {
*cell += k[k_idx] * delta[v_idx];
}
}
for (k_idx, row) in state[h].iter().enumerate() {
for (v_idx, cell) in row.iter().enumerate() {
core[h * dv + v_idx] += cell * q[k_idx] * scale;
}
}
}
let mut gated = vec![0f32; value_dim];
for h in 0..cfg.num_value_heads {
let base = h * dv;
let mean_sq = core[base..base + dv].iter().map(|x| x * x).sum::<f32>() / dv as f32;
let inv = 1.0 / (mean_sq + cfg.rms_norm_eps).sqrt();
for i in 0..dv {
gated[base + i] = core[base + i] * inv * raw.norm[i] * ref_silu(z[base + i]);
}
}
outputs.push(matvec(&raw.out, hidden, value_dim, &gated));
}
outputs
}
fn unequal_cfg() -> GdnConfig {
GdnConfig {
hidden_dim: 3,
num_key_heads: 1,
num_value_heads: 2,
key_head_dim: 2,
value_head_dim: 3,
conv_kernel_size: 3,
rms_norm_eps: 1e-5,
}
}
fn unequal_raw(cfg: &GdnConfig) -> RawGdn {
RawGdn {
qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
dt: vec![0.1, -0.05],
a: vec![-0.5, -0.75],
beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
norm: vec![1.1, 0.9, 1.3],
out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
}
}
#[test]
fn gdn_forward_token_tiny_dims_finite_and_shaped() {
let weights = make_weights();
let cfg = cfg();
let mut state = GdnState::new(&cfg);
let hidden = [0.2f32, -0.1, 0.3, -0.4];
let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
assert_eq!(out0.len(), HIDDEN);
assert!(out0.iter().all(|x| x.is_finite()));
let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
assert_eq!(out1.len(), HIDDEN);
assert!(out1.iter().all(|x| x.is_finite()));
assert!(
out0.iter()
.zip(out1.iter())
.any(|(a, b)| (a - b).abs() > 1e-6),
"recurrent state should change the second token"
);
}
#[test]
fn softplus_matches_closed_form_at_zero() {
assert!((softplus(0.0) - (2.0f32).ln()).abs() < 1e-6);
}
#[test]
fn unequal_head_geometry_matches_the_reference_split_and_replication() {
let cfg = unequal_cfg();
assert_eq!(cfg.key_dim(), 2, "key_dim = num_key_heads * key_head_dim");
assert_eq!(
cfg.value_dim(),
6,
"value_dim = num_value_heads * value_head_dim"
);
assert_eq!(
cfg.qkv_dim(),
10,
"conv_dim = 2*key_dim + value_dim; the equal-head formula would say 12"
);
assert_eq!(cfg.heads_per_key_group(), 2);
let raw = unequal_raw(&cfg);
let tokens = vec![
vec![0.2f32, -0.1, 0.3],
vec![-0.4f32, 0.25, 0.05],
vec![0.15f32, 0.35, -0.2],
];
let expected = reference_forward(&raw, &cfg, &tokens);
let weights = raw.to_weights(&cfg);
let mut state = GdnState::new(&cfg);
for (t, token) in tokens.iter().enumerate() {
let got = gdn_forward_token(&weights, &cfg, token, &mut state);
assert_eq!(got.len(), cfg.hidden_dim);
for (i, (g, e)) in got.iter().zip(expected[t].iter()).enumerate() {
assert!(
(g - e).abs() <= 1e-6 + 1e-5 * e.abs(),
"token {t} dim {i}: got {g}, reference {e}"
);
}
}
}
#[test]
fn replicating_one_key_head_equals_a_checkpoint_with_duplicated_key_rows() {
let shared = unequal_cfg(); let mut duplicated = shared;
duplicated.num_key_heads = 2;
let raw_shared = unequal_raw(&shared);
let hidden = shared.hidden_dim;
let dk = shared.key_head_dim;
let kernel = shared.conv_kernel_size;
let mut qkv_dup = Vec::with_capacity(duplicated.qkv_dim() * hidden);
let mut conv_dup = Vec::with_capacity(duplicated.qkv_dim() * kernel);
for part in 0..2 {
let w_src = part * shared.key_dim() * hidden;
let c_src = part * shared.key_dim() * kernel;
for _ in 0..2 {
qkv_dup.extend_from_slice(&raw_shared.qkv[w_src..w_src + dk * hidden]);
conv_dup.extend_from_slice(&raw_shared.conv[c_src..c_src + dk * kernel]);
}
}
qkv_dup.extend_from_slice(&raw_shared.qkv[2 * shared.key_dim() * hidden..]);
conv_dup.extend_from_slice(&raw_shared.conv[2 * shared.key_dim() * kernel..]);
let raw_dup = RawGdn {
qkv: qkv_dup,
conv: conv_dup,
gate: raw_shared.gate.clone(),
dt: raw_shared.dt.clone(),
a: raw_shared.a.clone(),
beta: raw_shared.beta.clone(),
alpha: raw_shared.alpha.clone(),
norm: raw_shared.norm.clone(),
out: raw_shared.out.clone(),
};
let w_shared = raw_shared.to_weights(&shared);
let w_dup = raw_dup.to_weights(&duplicated);
let mut s_shared = GdnState::new(&shared);
let mut s_dup = GdnState::new(&duplicated);
for token in [
vec![0.2f32, -0.1, 0.3],
vec![-0.4f32, 0.25, 0.05],
vec![0.15f32, 0.35, -0.2],
] {
let a = gdn_forward_token(&w_shared, &shared, &token, &mut s_shared);
let b = gdn_forward_token(&w_dup, &duplicated, &token, &mut s_dup);
for (x, y) in a.iter().zip(b.iter()) {
assert!((x - y).abs() < 1e-6, "{a:?} vs {b:?}");
}
}
}
#[test]
fn equal_head_geometry_stays_bit_identical_to_the_pre_generalization_output() {
const GOLDEN_STEP0: [u32; HIDDEN] = [1006633802, 998763940, 3163995192, 1006633802];
const GOLDEN_STEP1: [u32; HIDDEN] = [1006151492, 1000425832, 3164698040, 1006151492];
let weights = make_weights();
let cfg = cfg();
assert_eq!(cfg.qkv_dim(), QKV_DIM, "equal heads keep 3 * n_heads * dim");
assert_eq!(cfg.value_dim(), V_DIM);
assert_eq!(cfg.heads_per_key_group(), 1, "no replication when K == V");
let mut state = GdnState::new(&cfg);
let hidden = [0.2f32, -0.1, 0.3, -0.4];
let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
for (i, (got, want)) in out0.iter().zip(GOLDEN_STEP0.iter()).enumerate() {
assert_eq!(got.to_bits(), *want, "step 0 dim {i}: {got}");
}
for (i, (got, want)) in out1.iter().zip(GOLDEN_STEP1.iter()).enumerate() {
assert_eq!(got.to_bits(), *want, "step 1 dim {i}: {got}");
}
}
#[test]
fn recurrent_state_is_rectangular_when_key_and_value_head_dims_differ() {
let cfg = unequal_cfg();
let state = GdnState::new(&cfg);
assert_eq!(state.recurrent.len(), 2 * 3 * 2);
assert!(state.recurrent.iter().all(|x| *x == 0.0));
}
#[test]
#[should_panic(expected = "positive multiple")]
fn value_heads_not_a_multiple_of_key_heads_is_rejected_not_floored() {
let cfg = GdnConfig {
hidden_dim: 3,
num_key_heads: 3,
num_value_heads: 4,
key_head_dim: 2,
value_head_dim: 2,
conv_kernel_size: 2,
rms_norm_eps: 1e-5,
};
let raw = RawGdn {
qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
dt: vec![0.0; cfg.num_value_heads],
a: vec![-0.5; cfg.num_value_heads],
beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
norm: vec![1.0; cfg.value_head_dim],
out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
};
let weights = raw.to_weights(&cfg);
let mut state = GdnState::new(&cfg);
gdn_forward_token(&weights, &cfg, &[0.1, 0.2, 0.3], &mut state);
}
}