Skip to main content

ferrox_models/
gdn.rs

1//! Qwen-style Gated Delta Net (GDN) — linear-attention / SSM recurrent
2//! primitive for hybrid arches (`qwen35`, `qwen35moe`, `qwen3next`, …).
3//!
4//! Distinct from Kimi KDA (`kda.rs`): GDN uses a **fused** QKV projection,
5//! a **single** depthwise `ssm_conv1d` over the concatenated QKV channels,
6//! per-head `ssm_alpha` / `ssm_beta` gates, and decay
7//! `exp(softplus(α + ssm_dt) · ssm_a)` (GGUF `ssm_a` is typically
8//! `-exp(A_log)`). KDA is not a drop-in for this graph.
9//!
10//! ## GQA geometry (the shape rule this module implements)
11//!
12//! GDN is **GQA-shaped**: a real checkpoint has *fewer K heads than V
13//! heads* (`num_value_heads % num_key_heads == 0`) and the K and V head
14//! dims need not be equal (`key_head_dim ≠ value_head_dim` is legal).
15//! Transcribed from FreeToken's `qwen3_5_moe/gdn_reference.py`
16//! (`Qwen3_5GatedDeltaNetReference.forward`, no-cache path) and
17//! `LinearGatedDeltaGroupConfig`, where `num_key_heads` /
18//! `num_value_heads` / `key_head_dim` / `value_head_dim` are four
19//! independent numbers, not two:
20//!
21//! * **Split offsets.** The fused projection produces `conv_dim =
22//!   2·key_dim + value_dim` channels, with `key_dim = num_key_heads ·
23//!   key_head_dim` and `value_dim = num_value_heads · value_head_dim`,
24//!   and is split as `[key_dim, key_dim, value_dim]` — *not* into three
25//!   equal thirds. Assuming equality computes the K and V offsets wrong,
26//!   so **every** head reads a slice straddling the wrong tensor: the
27//!   layer stays finite and correctly shaped and silently returns
28//!   garbage rather than failing.
29//! * **Replication.** Each K head's q/k pair is replicated
30//!   `num_value_heads / num_key_heads` times (`repeat_interleave` on the
31//!   head axis) *before* the recurrence, so V head `h` reads K head
32//!   `h / rep`. Skipping the replication indexes q/k past the end of the
33//!   Q slice (into K, then into V) instead of reusing the shared head.
34//! * **Rectangular state.** The recurrent state is `[num_value_heads,
35//!   key_head_dim, value_head_dim]`, not square, and the q scale is
36//!   `key_head_dim^-0.5` (the *key* dim — the reference takes `dk` from
37//!   `key.shape[-1]`). Sizing the state from one head dim under-allocates
38//!   whenever `key_head_dim > value_head_dim` and mis-strides the
39//!   read-out in either direction.
40//!
41//! With `num_key_heads == num_value_heads` and `key_head_dim ==
42//! value_head_dim` every rule above collapses to the older equal-head
43//! path, bit for bit — pinned by
44//! `equal_head_geometry_stays_bit_identical_to_the_pre_generalization_output`.
45//!
46//! ## GGUF tensor name mapping (per layer `L`)
47//!
48//! | Role | GGUF name |
49//! |---|---|
50//! | Fused Q‖K‖V | `blk.{L}.attn_qkv.weight` |
51//! | Output / z gate | `blk.{L}.attn_gate.weight` |
52//! | Depthwise causal conv | `blk.{L}.ssm_conv1d.weight` |
53//! | Decay bias | `blk.{L}.ssm_dt.bias` (alt: `ssm_dt`) |
54//! | Decay scale | `blk.{L}.ssm_a` |
55//! | Input gate β | `blk.{L}.ssm_beta.weight` |
56//! | Forget raw α | `blk.{L}.ssm_alpha.weight` |
57//! | Output RMSNorm | `blk.{L}.ssm_norm.weight` |
58//! | Output projection | `blk.{L}.ssm_out.weight` |
59//!
60//! Legacy `qwen3next` may pack β/α into `ssm_ba` or fuse QKV+z into
61//! `ssm_in`; this module implements the split qwen35 layout only.
62//!
63//! GGUF weight load skeleton: [`crate::hybrid_gguf_loader`]. Serve still
64//! fail-closed — factory [`HybridEngine::reject`](crate::hybrid_engine::HybridEngine::reject).
65
66use ferrox_core::matmul::{rms_norm, silu};
67use ferrox_core::weight_matrix::WeightMatrix;
68
69/// Dims for one Qwen35-style GDN block, with **independent** K/V head
70/// counts and K/V head dims (see the module docs for why all four are
71/// separate numbers).
72///
73/// Invariant: `num_key_heads > 0` and `num_value_heads % num_key_heads ==
74/// 0`. The quotient is the q/k replication factor; a config that violated
75/// it would leave V heads with no K head to read from, so
76/// [`gdn_forward_token`] asserts it instead of flooring the division and
77/// silently pairing a V head with a K head that never fed it.
78#[derive(Debug, Clone, Copy)]
79pub struct GdnConfig {
80    pub hidden_dim: usize,
81    /// Number of Q/K heads in the checkpoint (`≤ num_value_heads`).
82    pub num_key_heads: usize,
83    /// Number of V heads — also the length of `ssm_dt` / `ssm_a` and the
84    /// row count of `ssm_beta` / `ssm_alpha`.
85    pub num_value_heads: usize,
86    /// Width of one Q/K head (`dk`; sets the `dk^-0.5` q scale).
87    pub key_head_dim: usize,
88    /// Width of one V head (`dv`; also the width of `ssm_norm`).
89    pub value_head_dim: usize,
90    pub conv_kernel_size: usize,
91    pub rms_norm_eps: f32,
92}
93
94impl GdnConfig {
95    /// Channels occupied by Q (and, separately, by K) in the fused
96    /// projection: `num_key_heads · key_head_dim`.
97    pub fn key_dim(&self) -> usize {
98        self.num_key_heads * self.key_head_dim
99    }
100
101    /// Channels occupied by V in the fused projection — also the width of
102    /// the `attn_gate` (z) projection and of `ssm_out`'s input.
103    pub fn value_dim(&self) -> usize {
104        self.num_value_heads * self.value_head_dim
105    }
106
107    /// Fused Q‖K‖V width: `key_dim + key_dim + value_dim`.
108    ///
109    /// This is the number the equal-head formula got wrong: `3 ·
110    /// num_value_heads · head_dim` only coincides with the real width
111    /// when both head counts *and* both head dims match. It is the same
112    /// arithmetic that yields the split offsets, so a wrong total means
113    /// wrong Q/K/V slices — not a length mismatch anyone would notice.
114    pub fn qkv_dim(&self) -> usize {
115        2 * self.key_dim() + self.value_dim()
116    }
117
118    /// q/k replication factor: how many V heads share one K head (the
119    /// `repeat_interleave` count in the reference). `1` for the
120    /// equal-head geometry.
121    pub fn heads_per_key_group(&self) -> usize {
122        self.num_value_heads / self.num_key_heads
123    }
124}
125
126/// Weights matching the qwen35 GGUF layout (see module docs).
127pub struct GdnWeights {
128    pub attn_qkv: WeightMatrix,  // [qkv_dim, hidden]
129    pub attn_gate: WeightMatrix, // [value_dim, hidden]
130    /// Depthwise taps, row-major `[qkv_dim, conv_kernel_size]`.
131    pub ssm_conv1d: Vec<f32>,
132    pub ssm_dt: Vec<f32>,        // [num_value_heads]
133    pub ssm_a: Vec<f32>,         // [num_value_heads]
134    pub ssm_beta: WeightMatrix,  // [num_value_heads, hidden]
135    pub ssm_alpha: WeightMatrix, // [num_value_heads, hidden]
136    pub ssm_norm: Vec<f32>,      // [value_head_dim]
137    pub ssm_out: WeightMatrix,   // [hidden, value_dim]
138}
139
140/// Fixed-size recurrent + short-conv state (unlike growing KV).
141pub struct GdnState {
142    conv_hist: Vec<f32>,
143    /// Flat `[num_value_heads, value_head_dim, key_head_dim]` — one
144    /// **rectangular** `state[v, k]` block per V head. The reference
145    /// stores the transpose (`[dk, dv]`); same elements, different
146    /// traversal order. Sizing this from a single `head_dim` truncates
147    /// the block whenever `key_head_dim != value_head_dim`, so the second
148    /// and later heads would read another head's memory.
149    recurrent: Vec<f32>,
150}
151
152impl GdnState {
153    pub fn new(cfg: &GdnConfig) -> Self {
154        Self {
155            conv_hist: Vec::new(),
156            recurrent: vec![0f32; cfg.num_value_heads * cfg.value_head_dim * cfg.key_head_dim],
157        }
158    }
159}
160
161fn softplus(x: f32) -> f32 {
162    if x > 20.0 {
163        x
164    } else {
165        (1.0 + x.exp()).ln()
166    }
167}
168
169fn sigmoid(x: f32) -> f32 {
170    1.0 / (1.0 + (-x).exp())
171}
172
173fn l2_normalize(v: &mut [f32], eps: f32) {
174    let norm_sq: f32 = v.iter().map(|x| x * x).sum();
175    let scale = 1.0 / (norm_sq + eps).sqrt();
176    for x in v.iter_mut() {
177        *x *= scale;
178    }
179}
180
181/// Depthwise causal conv over `dim` channels + SiLU (padding = kernel−1).
182fn causal_conv_step(
183    weight: &[f32],
184    history: &mut Vec<f32>,
185    current: &[f32],
186    kernel_size: usize,
187    dim: usize,
188) -> Vec<f32> {
189    let hist_len = history.len() / dim.max(1);
190    let missing = (kernel_size - 1).saturating_sub(hist_len);
191
192    let mut y = vec![0f32; dim];
193    for j in 0..kernel_size {
194        if j < missing {
195            continue;
196        }
197        let src: &[f32] = if j == kernel_size - 1 {
198            current
199        } else {
200            let hist_idx = j - missing;
201            &history[hist_idx * dim..(hist_idx + 1) * dim]
202        };
203        for d in 0..dim {
204            y[d] += weight[d * kernel_size + j] * src[d];
205        }
206    }
207    for v in y.iter_mut() {
208        *v = silu(*v);
209    }
210
211    history.extend_from_slice(current);
212    let max_hist_len = (kernel_size - 1) * dim;
213    if history.len() > max_hist_len {
214        let excess = history.len() - max_hist_len;
215        history.drain(0..excess);
216    }
217    y
218}
219
220/// One decode step of the GQA-shaped Gated Delta Net.
221///
222/// Handles `num_key_heads ≤ num_value_heads` and `key_head_dim ≠
223/// value_head_dim`: the fused projection is split at
224/// `[key_dim, key_dim, value_dim]`, V head `h` reads K head
225/// `h / heads_per_key_group()` (the reference's `repeat_interleave`), and
226/// each head's state is the rectangular `[value_head_dim, key_head_dim]`
227/// block. The equal-head geometry is the `heads_per_key_group() == 1`
228/// special case and runs identical arithmetic in identical order.
229pub fn gdn_forward_token(
230    weights: &GdnWeights,
231    cfg: &GdnConfig,
232    hidden: &[f32],
233    state: &mut GdnState,
234) -> Vec<f32> {
235    assert_eq!(hidden.len(), cfg.hidden_dim);
236    assert!(
237        cfg.num_key_heads > 0 && cfg.num_value_heads.is_multiple_of(cfg.num_key_heads),
238        "GDN needs num_value_heads ({}) to be a positive multiple of num_key_heads ({}); \
239         otherwise repeat_interleave has no whole replication factor and some V heads would \
240         silently read a K head that never fed them",
241        cfg.num_value_heads,
242        cfg.num_key_heads
243    );
244
245    let qkv_dim = cfg.qkv_dim();
246    let key_dim = cfg.key_dim();
247    let value_dim = cfg.value_dim();
248    let key_head_dim = cfg.key_head_dim;
249    let value_head_dim = cfg.value_head_dim;
250    let rep = cfg.heads_per_key_group();
251
252    let qkv_lin = weights.attn_qkv.apply(hidden);
253    let z = weights.attn_gate.apply(hidden);
254    let beta_raw = weights.ssm_beta.apply(hidden);
255    let alpha_raw = weights.ssm_alpha.apply(hidden);
256
257    let qkv = causal_conv_step(
258        &weights.ssm_conv1d,
259        &mut state.conv_hist,
260        &qkv_lin,
261        cfg.conv_kernel_size,
262        qkv_dim,
263    );
264
265    // torch.split(mixed_qkv, [key_dim, key_dim, value_dim], dim=-1).
266    let (q_all, rest) = qkv.split_at(key_dim);
267    let (k_all, v_all) = rest.split_at(key_dim);
268
269    let scale = 1.0 / (key_head_dim as f32).sqrt();
270    let mut y_flat = vec![0f32; value_dim];
271
272    #[allow(clippy::needless_range_loop)]
273    for h in 0..cfg.num_value_heads {
274        // repeat_interleave(rep, dim=head): V head h consumes K head h / rep.
275        let k_base = (h / rep) * key_head_dim;
276        let v_base = h * value_head_dim;
277        let mut q_h = q_all[k_base..k_base + key_head_dim].to_vec();
278        let mut k_h = k_all[k_base..k_base + key_head_dim].to_vec();
279        let v_h = &v_all[v_base..v_base + value_head_dim];
280
281        l2_normalize(&mut q_h, 1e-6);
282        l2_normalize(&mut k_h, 1e-6);
283        for x in q_h.iter_mut() {
284            *x *= scale;
285        }
286
287        // g = exp(softplus(α + dt) * A); A = ssm_a (often negative).
288        let gate = softplus(alpha_raw[h] + weights.ssm_dt[h]) * weights.ssm_a[h];
289        let decay = gate.exp();
290        let beta = sigmoid(beta_raw[h]);
291
292        let block = value_head_dim * key_head_dim;
293        let s_base = h * block;
294        let s = &mut state.recurrent[s_base..s_base + block];
295
296        // state *= decay
297        for cell in s.iter_mut() {
298            *cell *= decay;
299        }
300
301        // kv_mem[v] = sum_k state[v,k] * k[k]
302        let mut kv_mem = vec![0f32; value_head_dim];
303        for v_idx in 0..value_head_dim {
304            let mut acc = 0f32;
305            for k_idx in 0..key_head_dim {
306                acc += s[v_idx * key_head_dim + k_idx] * k_h[k_idx];
307            }
308            kv_mem[v_idx] = acc;
309        }
310
311        // state[v,k] += beta * (v - kv_mem)[v] * k[k]
312        for v_idx in 0..value_head_dim {
313            let delta = (v_h[v_idx] - kv_mem[v_idx]) * beta;
314            for k_idx in 0..key_head_dim {
315                s[v_idx * key_head_dim + k_idx] += delta * k_h[k_idx];
316            }
317        }
318
319        // y[v] = sum_k state[v,k] * q[k]
320        for v_idx in 0..value_head_dim {
321            let mut acc = 0f32;
322            for k_idx in 0..key_head_dim {
323                acc += s[v_idx * key_head_dim + k_idx] * q_h[k_idx];
324            }
325            y_flat[v_base + v_idx] = acc;
326        }
327    }
328
329    // Per-head RMSNorm on y (over value_head_dim), then SiLU(z) * normed.
330    let mut gated = vec![0f32; value_dim];
331    for h in 0..cfg.num_value_heads {
332        let base = h * value_head_dim;
333        let normed = rms_norm(
334            &y_flat[base..base + value_head_dim],
335            &weights.ssm_norm,
336            cfg.rms_norm_eps,
337        );
338        for i in 0..value_head_dim {
339            gated[base + i] = silu(z[base + i]) * normed[i];
340        }
341    }
342
343    weights.ssm_out.apply(&gated)
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349    use ferrox_core::tensor::Tensor;
350
351    const HIDDEN: usize = 4;
352    const N_HEADS: usize = 2;
353    const HEAD_DIM: usize = 2;
354    const CONV_K: usize = 2;
355    const QKV_DIM: usize = 3 * N_HEADS * HEAD_DIM; // 12
356    const V_DIM: usize = N_HEADS * HEAD_DIM; // 4
357
358    fn wm(data: &[f32], rows: usize, cols: usize) -> WeightMatrix {
359        assert_eq!(data.len(), rows * cols);
360        WeightMatrix::F32(Tensor::new(data.to_vec(), vec![rows, cols]))
361    }
362
363    fn cfg() -> GdnConfig {
364        GdnConfig {
365            hidden_dim: HIDDEN,
366            num_key_heads: N_HEADS,
367            num_value_heads: N_HEADS,
368            key_head_dim: HEAD_DIM,
369            value_head_dim: HEAD_DIM,
370            conv_kernel_size: CONV_K,
371            rms_norm_eps: 1e-5,
372        }
373    }
374
375    fn make_weights() -> GdnWeights {
376        // Deterministic tiny synthetic weights (not a golden oracle).
377        let mut qkv = Vec::with_capacity(QKV_DIM * HIDDEN);
378        for i in 0..QKV_DIM * HIDDEN {
379            qkv.push(((i % 7) as f32 - 3.0) * 0.1);
380        }
381        let mut gate = Vec::with_capacity(V_DIM * HIDDEN);
382        for i in 0..V_DIM * HIDDEN {
383            gate.push(((i % 5) as f32 - 2.0) * 0.08);
384        }
385        let mut conv = Vec::with_capacity(QKV_DIM * CONV_K);
386        for i in 0..QKV_DIM * CONV_K {
387            conv.push(if i % CONV_K == CONV_K - 1 { 1.0 } else { 0.1 });
388        }
389        let mut beta = Vec::with_capacity(N_HEADS * HIDDEN);
390        let mut alpha = Vec::with_capacity(N_HEADS * HIDDEN);
391        for i in 0..N_HEADS * HIDDEN {
392            beta.push(((i % 3) as f32 - 1.0) * 0.2);
393            alpha.push(((i % 4) as f32 - 1.5) * 0.15);
394        }
395        let mut out = Vec::with_capacity(HIDDEN * V_DIM);
396        for i in 0..HIDDEN * V_DIM {
397            out.push(((i % 6) as f32 - 2.5) * 0.12);
398        }
399        GdnWeights {
400            attn_qkv: wm(&qkv, QKV_DIM, HIDDEN),
401            attn_gate: wm(&gate, V_DIM, HIDDEN),
402            ssm_conv1d: conv,
403            ssm_dt: vec![0.1, -0.05],
404            // Negative A → decay ∈ (0, 1] after softplus·A + exp.
405            ssm_a: vec![-0.5, -0.75],
406            ssm_beta: wm(&beta, N_HEADS, HIDDEN),
407            ssm_alpha: wm(&alpha, N_HEADS, HIDDEN),
408            ssm_norm: vec![1.0, 1.0],
409            ssm_out: wm(&out, HIDDEN, V_DIM),
410        }
411    }
412
413    /// Raw (un-wrapped) tensor data for one GDN layer, so the reference
414    /// oracle below can run its own matvecs instead of borrowing this
415    /// module's [`WeightMatrix`] path.
416    struct RawGdn {
417        qkv: Vec<f32>,   // [qkv_dim, hidden]
418        gate: Vec<f32>,  // [value_dim, hidden]
419        conv: Vec<f32>,  // [qkv_dim, kernel]
420        dt: Vec<f32>,    // [num_value_heads]
421        a: Vec<f32>,     // [num_value_heads]
422        beta: Vec<f32>,  // [num_value_heads, hidden]
423        alpha: Vec<f32>, // [num_value_heads, hidden]
424        norm: Vec<f32>,  // [value_head_dim]
425        out: Vec<f32>,   // [hidden, value_dim]
426    }
427
428    impl RawGdn {
429        fn to_weights(&self, cfg: &GdnConfig) -> GdnWeights {
430            let h = cfg.hidden_dim;
431            GdnWeights {
432                attn_qkv: wm(&self.qkv, cfg.qkv_dim(), h),
433                attn_gate: wm(&self.gate, cfg.value_dim(), h),
434                ssm_conv1d: self.conv.clone(),
435                ssm_dt: self.dt.clone(),
436                ssm_a: self.a.clone(),
437                ssm_beta: wm(&self.beta, cfg.num_value_heads, h),
438                ssm_alpha: wm(&self.alpha, cfg.num_value_heads, h),
439                ssm_norm: self.norm.clone(),
440                ssm_out: wm(&self.out, h, cfg.value_dim()),
441            }
442        }
443    }
444
445    /// Deterministic spread of distinct, bounded, sign-alternating values:
446    /// no two channels of the fused projection carry the same number, so a
447    /// wrong split offset cannot accidentally alias onto a right answer.
448    fn fill(n: usize, seed: usize) -> Vec<f32> {
449        (0..n)
450            .map(|i| {
451                let k = (i * 37 + seed * 101) % 23;
452                (k as f32 - 11.0)
453                    * 0.043
454                    * if (i + seed).is_multiple_of(2) {
455                        1.0
456                    } else {
457                        -1.0
458                    }
459            })
460            .collect()
461    }
462
463    fn matvec(rows_data: &[f32], rows: usize, cols: usize, x: &[f32]) -> Vec<f32> {
464        assert_eq!(rows_data.len(), rows * cols);
465        assert_eq!(x.len(), cols);
466        (0..rows)
467            .map(|r| (0..cols).map(|c| rows_data[r * cols + c] * x[c]).sum())
468            .collect()
469    }
470
471    fn ref_sigmoid(x: f32) -> f32 {
472        1.0 / (1.0 + (-x).exp())
473    }
474
475    fn ref_softplus(x: f32) -> f32 {
476        (1.0 + x.exp()).ln()
477    }
478
479    fn ref_silu(x: f32) -> f32 {
480        x * ref_sigmoid(x)
481    }
482
483    fn ref_l2norm(v: &[f32]) -> Vec<f32> {
484        let sum_sq: f32 = v.iter().map(|x| x * x).sum();
485        let inv = 1.0 / (sum_sq + 1e-6).sqrt();
486        v.iter().map(|x| x * inv).collect()
487    }
488
489    /// Independent transcription of `gdn_reference.py`'s
490    /// `Qwen3_5GatedDeltaNetReference.forward` +
491    /// `recurrent_gated_delta_rule`, run over a whole token sequence in
492    /// the reference's own index order: the split is literally
493    /// `[key_dim, key_dim, value_dim]` (`:152`), q/k are materialized
494    /// through an explicit `repeat_interleave` (`:163`), and state is
495    /// `[dk, dv]` (this module stores the transpose). Nothing here calls
496    /// [`gdn_forward_token`], so it is an oracle rather than a restatement
497    /// of the code under test.
498    ///
499    /// `ssm_a` follows the GGUF convention this module documents
500    /// (`ssm_a == -exp(A_log)`), so `g = softplus(α + dt) · ssm_a` spells
501    /// the reference's `-A_log.exp() * softplus(a + dt_bias)`.
502    fn reference_forward(raw: &RawGdn, cfg: &GdnConfig, tokens: &[Vec<f32>]) -> Vec<Vec<f32>> {
503        let hidden = cfg.hidden_dim;
504        let qkv_dim = cfg.qkv_dim();
505        let key_dim = cfg.key_dim();
506        let value_dim = cfg.value_dim();
507        let dk = cfg.key_head_dim;
508        let dv = cfg.value_head_dim;
509        let kernel = cfg.conv_kernel_size;
510        let rep = cfg.num_value_heads / cfg.num_key_heads;
511        let scale = 1.0 / (dk as f32).sqrt();
512
513        // in_proj_qkv over the sequence, then the depthwise causal conv
514        // with padding = kernel-1 (positions before 0 read as zero) + silu.
515        let mixed: Vec<Vec<f32>> = tokens
516            .iter()
517            .map(|h| matvec(&raw.qkv, qkv_dim, hidden, h))
518            .collect();
519        let mut conved = vec![vec![0f32; qkv_dim]; tokens.len()];
520        for (t, conved_t) in conved.iter_mut().enumerate() {
521            for (d, out_d) in conved_t.iter_mut().enumerate() {
522                let mut acc = 0f32;
523                for j in 0..kernel {
524                    let src = t as isize - (kernel as isize - 1) + j as isize;
525                    if src < 0 {
526                        continue;
527                    }
528                    acc += raw.conv[d * kernel + j] * mixed[src as usize][d];
529                }
530                *out_d = ref_silu(acc);
531            }
532        }
533
534        // state[h][k][v] — the reference's [num_v_heads, dk, dv] layout.
535        let mut state = vec![vec![vec![0f32; dv]; dk]; cfg.num_value_heads];
536        let mut outputs = Vec::with_capacity(tokens.len());
537
538        for (t, token) in tokens.iter().enumerate() {
539            let z = matvec(&raw.gate, value_dim, hidden, token);
540            let a_raw = matvec(&raw.alpha, cfg.num_value_heads, hidden, token);
541            let b_raw = matvec(&raw.beta, cfg.num_value_heads, hidden, token);
542
543            let q_slice = &conved[t][0..key_dim];
544            let k_slice = &conved[t][key_dim..2 * key_dim];
545            let v_slice = &conved[t][2 * key_dim..];
546
547            // repeat_interleave(rep) on the head axis of q and k.
548            let mut q_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
549            let mut k_heads: Vec<Vec<f32>> = Vec::with_capacity(cfg.num_value_heads);
550            for kh in 0..cfg.num_key_heads {
551                for _ in 0..rep {
552                    q_heads.push(q_slice[kh * dk..(kh + 1) * dk].to_vec());
553                    k_heads.push(k_slice[kh * dk..(kh + 1) * dk].to_vec());
554                }
555            }
556
557            let mut core = vec![0f32; value_dim];
558            for h in 0..cfg.num_value_heads {
559                let q = ref_l2norm(&q_heads[h]);
560                let k = ref_l2norm(&k_heads[h]);
561                let v = &v_slice[h * dv..(h + 1) * dv];
562
563                let decay = (ref_softplus(a_raw[h] + raw.dt[h]) * raw.a[h]).exp();
564                let beta = ref_sigmoid(b_raw[h]);
565
566                for row in state[h].iter_mut() {
567                    for cell in row.iter_mut() {
568                        *cell *= decay;
569                    }
570                }
571                // kv_mem = (state * k[:, None]).sum(dim=-2)
572                let mut kv_mem = vec![0f32; dv];
573                for (k_idx, row) in state[h].iter().enumerate() {
574                    for (v_idx, cell) in row.iter().enumerate() {
575                        kv_mem[v_idx] += cell * k[k_idx];
576                    }
577                }
578                // state += k[:, None] * ((v - kv_mem) * beta)[None, :]
579                let delta: Vec<f32> = (0..dv).map(|i| (v[i] - kv_mem[i]) * beta).collect();
580                for (k_idx, row) in state[h].iter_mut().enumerate() {
581                    for (v_idx, cell) in row.iter_mut().enumerate() {
582                        *cell += k[k_idx] * delta[v_idx];
583                    }
584                }
585                // out = (state * (q * scale)[:, None]).sum(dim=-2)
586                for (k_idx, row) in state[h].iter().enumerate() {
587                    for (v_idx, cell) in row.iter().enumerate() {
588                        core[h * dv + v_idx] += cell * q[k_idx] * scale;
589                    }
590                }
591            }
592
593            // RMSNormGated over head_v_dim groups, norm_before_gate=True.
594            let mut gated = vec![0f32; value_dim];
595            for h in 0..cfg.num_value_heads {
596                let base = h * dv;
597                let mean_sq = core[base..base + dv].iter().map(|x| x * x).sum::<f32>() / dv as f32;
598                let inv = 1.0 / (mean_sq + cfg.rms_norm_eps).sqrt();
599                for i in 0..dv {
600                    gated[base + i] = core[base + i] * inv * raw.norm[i] * ref_silu(z[base + i]);
601                }
602            }
603            outputs.push(matvec(&raw.out, hidden, value_dim, &gated));
604        }
605        outputs
606    }
607
608    /// A GQA-shaped geometry: one K head feeding two V heads,
609    /// `key_head_dim` 2 against `value_head_dim` 3. `qkv_dim` is
610    /// `2·2 + 6 = 10` and the split offsets are 0 / 2 / 4.
611    fn unequal_cfg() -> GdnConfig {
612        GdnConfig {
613            hidden_dim: 3,
614            num_key_heads: 1,
615            num_value_heads: 2,
616            key_head_dim: 2,
617            value_head_dim: 3,
618            conv_kernel_size: 3,
619            rms_norm_eps: 1e-5,
620        }
621    }
622
623    fn unequal_raw(cfg: &GdnConfig) -> RawGdn {
624        RawGdn {
625            qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
626            gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
627            conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
628            dt: vec![0.1, -0.05],
629            a: vec![-0.5, -0.75],
630            beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
631            alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
632            norm: vec![1.1, 0.9, 1.3],
633            out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
634        }
635    }
636
637    #[test]
638    fn gdn_forward_token_tiny_dims_finite_and_shaped() {
639        let weights = make_weights();
640        let cfg = cfg();
641        let mut state = GdnState::new(&cfg);
642        let hidden = [0.2f32, -0.1, 0.3, -0.4];
643
644        let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
645        assert_eq!(out0.len(), HIDDEN);
646        assert!(out0.iter().all(|x| x.is_finite()));
647
648        let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
649        assert_eq!(out1.len(), HIDDEN);
650        assert!(out1.iter().all(|x| x.is_finite()));
651        // Second step must see non-zero recurrent state → different output.
652        assert!(
653            out0.iter()
654                .zip(out1.iter())
655                .any(|(a, b)| (a - b).abs() > 1e-6),
656            "recurrent state should change the second token"
657        );
658    }
659
660    #[test]
661    fn softplus_matches_closed_form_at_zero() {
662        assert!((softplus(0.0) - (2.0f32).ln()).abs() < 1e-6);
663    }
664
665    /// **The central test — it fails against the pre-generalization
666    /// implementation.** The geometry is one K head / two V heads with
667    /// `key_head_dim = 2` and `value_head_dim = 3`, so the fused
668    /// projection is `2·2 + 2·3 = 10` wide and splits at `[2, 2, 6]`. The
669    /// equal-head formula (`3 · num_v_heads · head_dim`) claims 12 wide
670    /// splitting at `[4, 4, 4]`: it reads Q/K/V from offsets that belong
671    /// to the neighbouring tensor, and no single `head_dim` rescues it,
672    /// because K and V head dims genuinely differ here.
673    ///
674    /// Expected values come from [`reference_forward`] — a transcription
675    /// of `gdn_reference.py:152` (the `[key_dim, key_dim, value_dim]`
676    /// split) and `:163` (`repeat_interleave` of q/k up to `num_v_heads`)
677    /// in the reference's own `[dk, dv]` state layout — **not** from this
678    /// module's own output.
679    #[test]
680    fn unequal_head_geometry_matches_the_reference_split_and_replication() {
681        let cfg = unequal_cfg();
682        // Offsets straight from the reference formula, spelled out.
683        assert_eq!(cfg.key_dim(), 2, "key_dim = num_key_heads * key_head_dim");
684        assert_eq!(
685            cfg.value_dim(),
686            6,
687            "value_dim = num_value_heads * value_head_dim"
688        );
689        assert_eq!(
690            cfg.qkv_dim(),
691            10,
692            "conv_dim = 2*key_dim + value_dim; the equal-head formula would say 12"
693        );
694        assert_eq!(cfg.heads_per_key_group(), 2);
695
696        let raw = unequal_raw(&cfg);
697        let tokens = vec![
698            vec![0.2f32, -0.1, 0.3],
699            vec![-0.4f32, 0.25, 0.05],
700            vec![0.15f32, 0.35, -0.2],
701        ];
702        let expected = reference_forward(&raw, &cfg, &tokens);
703
704        let weights = raw.to_weights(&cfg);
705        let mut state = GdnState::new(&cfg);
706        for (t, token) in tokens.iter().enumerate() {
707            let got = gdn_forward_token(&weights, &cfg, token, &mut state);
708            assert_eq!(got.len(), cfg.hidden_dim);
709            for (i, (g, e)) in got.iter().zip(expected[t].iter()).enumerate() {
710                assert!(
711                    (g - e).abs() <= 1e-6 + 1e-5 * e.abs(),
712                    "token {t} dim {i}: got {g}, reference {e}"
713                );
714            }
715        }
716    }
717
718    /// Replicating one K head across `rep` V heads must be *exactly* the
719    /// same computation as a checkpoint that stored those `rep` K rows
720    /// duplicated on disk. If the replication indexed `h * key_head_dim`
721    /// instead of `(h / rep) * key_head_dim`, this equivalence breaks —
722    /// and for the one-K-head config it would read past the Q slice into
723    /// K, which no shape check would catch.
724    #[test]
725    fn replicating_one_key_head_equals_a_checkpoint_with_duplicated_key_rows() {
726        let shared = unequal_cfg(); // 1 K head → 2 V heads
727        let mut duplicated = shared;
728        duplicated.num_key_heads = 2; // same math, K/Q rows stored twice
729
730        let raw_shared = unequal_raw(&shared);
731        let hidden = shared.hidden_dim;
732        let dk = shared.key_head_dim;
733        let kernel = shared.conv_kernel_size;
734
735        // Rebuild the fused projection (and its conv taps) with the single
736        // K head's Q and K rows physically duplicated; V rows untouched.
737        let mut qkv_dup = Vec::with_capacity(duplicated.qkv_dim() * hidden);
738        let mut conv_dup = Vec::with_capacity(duplicated.qkv_dim() * kernel);
739        for part in 0..2 {
740            // Q block, then K block.
741            let w_src = part * shared.key_dim() * hidden;
742            let c_src = part * shared.key_dim() * kernel;
743            for _ in 0..2 {
744                qkv_dup.extend_from_slice(&raw_shared.qkv[w_src..w_src + dk * hidden]);
745                conv_dup.extend_from_slice(&raw_shared.conv[c_src..c_src + dk * kernel]);
746            }
747        }
748        qkv_dup.extend_from_slice(&raw_shared.qkv[2 * shared.key_dim() * hidden..]);
749        conv_dup.extend_from_slice(&raw_shared.conv[2 * shared.key_dim() * kernel..]);
750
751        let raw_dup = RawGdn {
752            qkv: qkv_dup,
753            conv: conv_dup,
754            gate: raw_shared.gate.clone(),
755            dt: raw_shared.dt.clone(),
756            a: raw_shared.a.clone(),
757            beta: raw_shared.beta.clone(),
758            alpha: raw_shared.alpha.clone(),
759            norm: raw_shared.norm.clone(),
760            out: raw_shared.out.clone(),
761        };
762
763        let w_shared = raw_shared.to_weights(&shared);
764        let w_dup = raw_dup.to_weights(&duplicated);
765        let mut s_shared = GdnState::new(&shared);
766        let mut s_dup = GdnState::new(&duplicated);
767        for token in [
768            vec![0.2f32, -0.1, 0.3],
769            vec![-0.4f32, 0.25, 0.05],
770            vec![0.15f32, 0.35, -0.2],
771        ] {
772            let a = gdn_forward_token(&w_shared, &shared, &token, &mut s_shared);
773            let b = gdn_forward_token(&w_dup, &duplicated, &token, &mut s_dup);
774            for (x, y) in a.iter().zip(b.iter()) {
775                assert!((x - y).abs() < 1e-6, "{a:?} vs {b:?}");
776            }
777        }
778    }
779
780    /// The regression this generalization must not break: with
781    /// `num_key_heads == num_value_heads` and `key_head_dim ==
782    /// value_head_dim` the layer must return **bit-identical** floats to
783    /// the pre-generalization equal-head implementation. The constants are
784    /// the raw `f32` bit patterns that implementation produced on
785    /// [`make_weights`] / [`cfg`], so a reordered accumulation, a scale
786    /// taken from the wrong head dim, or an altered state stride shows up
787    /// here instead of silently shifting every equal-head checkpoint's
788    /// output.
789    #[test]
790    fn equal_head_geometry_stays_bit_identical_to_the_pre_generalization_output() {
791        const GOLDEN_STEP0: [u32; HIDDEN] = [1006633802, 998763940, 3163995192, 1006633802];
792        const GOLDEN_STEP1: [u32; HIDDEN] = [1006151492, 1000425832, 3164698040, 1006151492];
793
794        let weights = make_weights();
795        let cfg = cfg();
796        assert_eq!(cfg.qkv_dim(), QKV_DIM, "equal heads keep 3 * n_heads * dim");
797        assert_eq!(cfg.value_dim(), V_DIM);
798        assert_eq!(cfg.heads_per_key_group(), 1, "no replication when K == V");
799
800        let mut state = GdnState::new(&cfg);
801        let hidden = [0.2f32, -0.1, 0.3, -0.4];
802        let out0 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
803        let out1 = gdn_forward_token(&weights, &cfg, &hidden, &mut state);
804
805        for (i, (got, want)) in out0.iter().zip(GOLDEN_STEP0.iter()).enumerate() {
806            assert_eq!(got.to_bits(), *want, "step 0 dim {i}: {got}");
807        }
808        for (i, (got, want)) in out1.iter().zip(GOLDEN_STEP1.iter()).enumerate() {
809            assert_eq!(got.to_bits(), *want, "step 1 dim {i}: {got}");
810        }
811    }
812
813    /// The recurrent state is `[num_value_heads, value_head_dim,
814    /// key_head_dim]`. A square `head_dim × head_dim` block would
815    /// allocate 2·2·2 = 8 floats here instead of 2·3·2 = 12, and the
816    /// second head's read-out would run off the end of the buffer.
817    #[test]
818    fn recurrent_state_is_rectangular_when_key_and_value_head_dims_differ() {
819        let cfg = unequal_cfg();
820        let state = GdnState::new(&cfg);
821        assert_eq!(state.recurrent.len(), 2 * 3 * 2);
822        assert!(state.recurrent.iter().all(|x| *x == 0.0));
823    }
824
825    #[test]
826    #[should_panic(expected = "positive multiple")]
827    fn value_heads_not_a_multiple_of_key_heads_is_rejected_not_floored() {
828        let cfg = GdnConfig {
829            hidden_dim: 3,
830            num_key_heads: 3,
831            num_value_heads: 4,
832            key_head_dim: 2,
833            value_head_dim: 2,
834            conv_kernel_size: 2,
835            rms_norm_eps: 1e-5,
836        };
837        let raw = RawGdn {
838            qkv: fill(cfg.qkv_dim() * cfg.hidden_dim, 1),
839            gate: fill(cfg.value_dim() * cfg.hidden_dim, 2),
840            conv: fill(cfg.qkv_dim() * cfg.conv_kernel_size, 3),
841            dt: vec![0.0; cfg.num_value_heads],
842            a: vec![-0.5; cfg.num_value_heads],
843            beta: fill(cfg.num_value_heads * cfg.hidden_dim, 4),
844            alpha: fill(cfg.num_value_heads * cfg.hidden_dim, 5),
845            norm: vec![1.0; cfg.value_head_dim],
846            out: fill(cfg.hidden_dim * cfg.value_dim(), 6),
847        };
848        let weights = raw.to_weights(&cfg);
849        let mut state = GdnState::new(&cfg);
850        gdn_forward_token(&weights, &cfg, &[0.1, 0.2, 0.3], &mut state);
851    }
852}