Skip to main content

cortiq_engine/
linear_core.rs

1//! Linear-attention cores, selected by `arch.linear_core.kind`
2//! (descriptor-driven operators — Patent 15 claim 8).
3//!
4//! Two tracks (owner decision 2026-07-04):
5//!
6//! * `gated_delta_net` — the faithful vendor operator (Qwen3.5 /
7//!   Qwen3-Next). Default for models that ship GDN weights: conversion
8//!   carries the tensors 1:1 and needs no training. Port of the
9//!   validated `gated_delta_net` (vmfcore/rust/src/forward.rs) against
10//!   the numpy/torch oracle (vmfcore/gdn_layer.py).
11//!
12//! * `vmf_phase` — the legacy additive core: token carries a phase θ; kernel
13//!   φ(θ) = [cos θ; sin θ] gives a linear factorization; the recurrent
14//!   state S[head][p2, dv] uses decay exp(−exp(A_log)).
15//!   Noise-robust and simpler than vendor recurrences. Exotic operators
16//!   are folded onto it at CONVERT time (`--linear-core vmf_phase`) and
17//!   quality is restored by the offline heal — the research track and
18//!   the production mechanism for Patent-15 skills (mask→heal→compress).
19//!
20//! * `vmf_phase_delta_v1` — the operator-tagged normalized in-place
21//!   Phase-Delta variant: decay first, read the old value under k, write the
22//!   gated residual, then read under q. It shares all VMF tensors/state but
23//!   is selected per layer by the CMF linear-core record.
24//!
25//! Both cores implement the same contract: `*_forward` (one position,
26//! advances the state) and `*_pair` (fused two positions; lane 1
27//! commits, lane 2 is tentative in `scratch` for speculative verify).
28//! State lives in the layer's `linear_state: Vec<f32>` and is resized
29//! lazily by the core itself.
30
31use crate::pool::Pool;
32use crate::qtensor::QTensor;
33use cortiq_core::TensorDtype;
34use std::sync::OnceLock;
35use std::sync::atomic::{AtomicU64, Ordering};
36
37static PERF_GDN_FORWARD_CALLS: AtomicU64 = AtomicU64::new(0);
38static PERF_GDN_FORWARD_NS: AtomicU64 = AtomicU64::new(0);
39static PERF_GDN_BATCH_CALLS: AtomicU64 = AtomicU64::new(0);
40static PERF_GDN_BATCH_NS: AtomicU64 = AtomicU64::new(0);
41static PERF_GDN_STEP_CALLS: AtomicU64 = AtomicU64::new(0);
42static PERF_GDN_STEP_NS: AtomicU64 = AtomicU64::new(0);
43
44fn perf_enabled() -> bool {
45    static ON: OnceLock<bool> = OnceLock::new();
46    *ON.get_or_init(|| std::env::var("CMF_PERF_PROFILE").as_deref() == Ok("1"))
47}
48
49/// Weights of one vmf_phase layer (`model.layers.{i}.vmf_attn.*`).
50pub struct VmfPhaseWeights {
51    /// [nh·nphase, hidden] — query phase projection
52    pub thq: QTensor,
53    /// [nh·nphase, hidden] — key phase projection
54    pub thk: QTensor,
55    /// [nh·dv, hidden]
56    pub v_proj: QTensor,
57    /// [hidden, nh·dv]
58    pub out_proj: QTensor,
59    /// Per-component decay exp(−exp(A_log)), len nh·2·nphase (precomputed).
60    pub decay: Vec<f64>,
61    /// Short causal depthwise conv before the projections (embryo genomes
62    /// born with `--conv-k`): flat `[hidden·k]` taps, identity-initialised
63    /// at training start. The projections were TRAINED on the conv output —
64    /// running without it scrambles the layer (measured: train-val ppl 60
65    /// exported to runtime ppl 93 539). None = pre-conv genomes, untouched.
66    pub conv: Option<Vec<f32>>,
67    /// Selective-write input gate κ (hybrid_k core, stage 71): weight
68    /// [nh, hidden] + bias [nh]; κ_h = σ(W_k·x + b)_h multiplies the
69    /// state WRITE (S = decay·S + κ·φk⊗v). None = classic phase core,
70    /// bit-identical to the pre-κ kernel. Measured at mechanism level:
71    /// knee ×2–6 earlier, restores correlated-noise robustness, LM
72    /// crossover vs softmax at SEQ 512 (experiments/lc_final_merged.json).
73    pub k_gate: Option<(QTensor, Vec<f32>)>,
74    /// Select the normalized in-place Phase-Delta recurrence. This bit is
75    /// derived from the versioned CMF selector and is never inferred from
76    /// tensor presence, so legacy `vmf_phase` files retain their exact
77    /// additive semantics.
78    pub phase_delta: bool,
79}
80
81#[derive(Clone, Copy)]
82pub struct VmfPhaseCfg {
83    pub num_heads: usize,
84    pub nphase: usize,
85    pub value_head_dim: usize,
86    pub hidden_size: usize,
87    /// Phase-mass correction: scales the phase toward zero —
88    /// θ_eff = θ/(1+mass) — which widens the phase kernel.
89    /// Measured (experiments/vmf_native_core*.py) to restore noise
90    /// robustness when the phase projection is FIXED (exactly CMF's
91    /// fold-before-heal regime: thq/thk are init, not trained) — recall
92    /// 3%→91% at moderate noise; redundant once the projection is
93    /// healed. 0.0 = massless Goldstone (bit-identical to prior kernel).
94    /// Set via CMF_PHASE_MASS. Validated at mechanism level, not yet LM.
95    pub phase_mass: f32,
96}
97
98impl VmfPhaseCfg {
99    pub fn state_len(&self) -> usize {
100        self.num_heads * 2 * self.nphase * self.value_head_dim
101    }
102}
103
104/// One normalized in-place Phase-Delta step.  The state is feature-major
105/// `[2*nphase, value_head_dim]`; `state` may carry a convolution ring after
106/// that prefix, so this helper only touches the recurrent matrix.  Keeping
107/// the decay/read/write/read ordering literal is what gives beta=1 and
108/// q=k the exact-overwrite witness from the trainer contract.
109fn phase_delta_step_f32(
110    thq: &[f32],
111    thk: &[f32],
112    v: &[f32],
113    decay: &[f64],
114    kap: Option<&[f32]>,
115    cfg: &VmfPhaseCfg,
116    state: &mut [f32],
117    out: &mut [f32],
118) {
119    let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
120    let p2 = 2 * nph;
121    let scale = 1.0f64 / (nph as f64).sqrt();
122    debug_assert_eq!(thq.len(), nh * nph);
123    debug_assert_eq!(thk.len(), nh * nph);
124    debug_assert_eq!(v.len(), nh * dv);
125    debug_assert_eq!(decay.len(), nh * p2);
126    debug_assert!(state.len() >= nh * p2 * dv);
127    debug_assert!(out.len() >= nh * dv);
128
129    // Reuse one head's workspace. The old-value reduction must finish before
130    // any cell is overwritten, but neither it nor the phase features needs a
131    // fresh allocation per head.
132    let mut r = vec![0.0f64; dv];
133    let mut key = vec![0.0f64; p2];
134    for h in 0..nh {
135        let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
136        let thk_h = &thk[h * nph..(h + 1) * nph];
137        let thq_h = &thq[h * nph..(h + 1) * nph];
138        let vt = &v[h * dv..(h + 1) * dv];
139        let ot = &mut out[h * dv..(h + 1) * dv];
140        let dec = &decay[h * p2..(h + 1) * p2];
141        let beta = kap.map_or(1.0f64, |k| k[h] as f64);
142
143        // P = D_gamma S_prev; r = k^T P.  P is kept as f64 only for this
144        // token's arithmetic, while the recurrent cells remain f32 exactly
145        // like the established runtime VMF path.
146        r.fill(0.0);
147        for f in 0..p2 {
148            let kf = if f < nph {
149                scale * (thk_h[f] as f64).cos()
150            } else {
151                scale * (thk_h[f - nph] as f64).sin()
152            };
153            key[f] = kf;
154            let row = &s[f * dv..(f + 1) * dv];
155            for d in 0..dv {
156                r[d] += kf * (dec[f] * row[d] as f64);
157            }
158        }
159
160        // S = P + beta*k*(v-r), then o = q^T S (post-write).
161        // On ARM fuse the write and read of each row: no third state sweep.
162        // On x86 keep separate loops: mixing f64 updates and f32 reductions
163        // prevents vectorization and regressed the measured Xeon kernel.
164        // Both paths avoid reevaluating the key's transcendental functions.
165        // Crucially read the rounded f32 cell, NOT its f64 precursor. Keep
166        // feature accumulation order and the f32 add exactly as before.
167        for f in 0..p2 {
168            let kf = key[f];
169            #[cfg(target_arch = "aarch64")]
170            let qf = if f < nph {
171                scale * (thq_h[f] as f64).cos()
172            } else {
173                scale * (thq_h[f - nph] as f64).sin()
174            };
175            let row = &mut s[f * dv..(f + 1) * dv];
176            for d in 0..dv {
177                let p = dec[f] * row[d] as f64;
178                row[d] = (p + beta * kf * (vt[d] as f64 - r[d])) as f32;
179                #[cfg(target_arch = "aarch64")]
180                {
181                    ot[d] += (qf * row[d] as f64) as f32;
182                }
183            }
184        }
185        #[cfg(not(target_arch = "aarch64"))]
186        for f in 0..p2 {
187            let qf = if f < nph {
188                scale * (thq_h[f] as f64).cos()
189            } else {
190                scale * (thq_h[f - nph] as f64).sin()
191            };
192            let row = &s[f * dv..(f + 1) * dv];
193            for d in 0..dv {
194                ot[d] += (qf * row[d] as f64) as f32;
195            }
196        }
197    }
198}
199
200/// One recurrent step for one head-set given projected phases/values.
201/// `state` is S[nh][p2, dv] stored f32 (per-element math in f64 — the
202/// storage halves, each step's arithmetic keeps the old precision).
203fn phase_step(
204    thq: &[f32],
205    thk: &[f32],
206    v: &[f32],
207    decay: &[f64],
208    kap: Option<&[f32]>,
209    cfg: &VmfPhaseCfg,
210    state: &mut [f32],
211    out: &mut [f32],
212    phase_delta: bool,
213) {
214    let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
215    if phase_delta {
216        phase_delta_step_f32(thq, thk, v, decay, kap, cfg, state, out);
217        return;
218    }
219    // Phase-mass correction: θ_eff = θ/(1+mass). mass=0 → factor 1 → no-op.
220    let mscale = 1.0f64 / (1.0 + cfg.phase_mass as f64);
221    let p2 = 2 * nph;
222    for h in 0..nh {
223        let s = &mut state[h * p2 * dv..(h + 1) * p2 * dv];
224        let thk_h = &thk[h * nph..(h + 1) * nph];
225        let thq_h = &thq[h * nph..(h + 1) * nph];
226        let vt = &v[h * dv..(h + 1) * dv];
227        let ot = &mut out[h * dv..(h + 1) * dv];
228        let dec = &decay[h * p2..(h + 1) * p2];
229        // Selective write (hybrid_k): κ scales what enters the recurrent state.
230        let kh = kap.map_or(1.0f64, |k| k[h] as f64);
231        for f in 0..p2 {
232            // φ(θ) = [cos·nph, sin·nph], θ scaled by the correction factor.
233            let (fk, fq) = if f < nph {
234                (
235                    (thk_h[f] as f64 * mscale).cos(),
236                    (thq_h[f] as f64 * mscale).cos(),
237                )
238            } else {
239                (
240                    (thk_h[f - nph] as f64 * mscale).sin(),
241                    (thq_h[f - nph] as f64 * mscale).sin(),
242                )
243            };
244            let fkw = fk * kh;
245            let row = &mut s[f * dv..(f + 1) * dv];
246            let dcf = dec[f];
247            for d in 0..dv {
248                // S = decay·S + κ·φk⊗v (f64 math, f32 cell)
249                let cell = dcf * row[d] as f64 + fkw * vt[d] as f64;
250                row[d] = cell as f32;
251                ot[d] += (fq * cell) as f32; // o = Σ φq·S
252            }
253        }
254    }
255}
256
257/// κ_h = σ(W_k·x + b)_h — the per-head write gate (None when the layer
258/// has no k_gate tensors: classic phase core).
259fn kappa_of(x: &[f32], w: &VmfPhaseWeights, nh: usize, pool: Option<&Pool>) -> Option<Vec<f32>> {
260    let (kw, kb) = w.k_gate.as_ref()?;
261    let mut k = vec![0.0f32; nh];
262    kw.matvec(x, &mut k, pool);
263    for (v, b) in k.iter_mut().zip(kb) {
264        *v = 1.0 / (1.0 + (-(*v + b)).exp());
265    }
266    Some(k)
267}
268
269/// Causal depthwise conv over the mixer input, with the last k−1 inputs
270/// ringed at the TAIL of the layer state (oldest first). Returns the
271/// convolved input; a layer without conv taps passes through untouched.
272fn conv_in(
273    x: &[f32],
274    w: &VmfPhaseWeights,
275    cfg: &VmfPhaseCfg,
276    state: &mut Vec<f32>,
277) -> Option<Vec<f32>> {
278    let taps = w.conv.as_ref()?;
279    let h = cfg.hidden_size;
280    let k = taps.len() / h.max(1);
281    if k < 2 || taps.len() != h * k {
282        return None;
283    }
284    let ring = (k - 1) * h;
285    let base = cfg.state_len();
286    if state.len() != base + ring {
287        // The phase part resets alongside — a fresh sequence either way.
288        let mut ns = vec![0f32; base + ring];
289        let n = state.len().min(base);
290        ns[..n].copy_from_slice(&state[..n]);
291        *state = ns;
292    }
293    let mut y = vec![0.0f32; h];
294    for c in 0..h {
295        // taps j = 0..k−2 read the ring (oldest first), tap k−1 reads x.
296        let mut acc = taps[c * k + k - 1] * x[c];
297        for j in 0..k - 1 {
298            acc += taps[c * k + j] * state[base + j * h + c];
299        }
300        y[c] = acc;
301    }
302    conv_ring_push(x, h, base, state);
303    Some(y)
304}
305
306/// Rotate the conv ring at `base`: drop the oldest input, append `x`.
307fn conv_ring_push(x: &[f32], h: usize, base: usize, state: &mut [f32]) {
308    let ring = state.len() - base;
309    state.copy_within(base + h.., base);
310    let at = base + ring - h;
311    state[at..at + h].copy_from_slice(&x[..h]);
312}
313
314/// Forward one position through a vmf_phase layer, advancing `state`.
315pub fn vmf_phase_forward(
316    x: &[f32],
317    w: &VmfPhaseWeights,
318    cfg: &VmfPhaseCfg,
319    state: &mut Vec<f32>,
320    pool: Option<&Pool>,
321) -> Vec<f32> {
322    if w.conv.is_none() && state.len() != cfg.state_len() {
323        *state = vec![0f32; cfg.state_len()];
324    }
325    let xc = conv_in(x, w, cfg, state);
326    let x = xc.as_deref().unwrap_or(x);
327    let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
328
329    let mut thq = vec![0.0f32; nh * nph];
330    w.thq.matvec(x, &mut thq, pool);
331    let mut thk = vec![0.0f32; nh * nph];
332    w.thk.matvec(x, &mut thk, pool);
333    let mut v = vec![0.0f32; nh * dv];
334    w.v_proj.matvec(x, &mut v, pool);
335
336    let kap = kappa_of(x, w, nh, pool);
337    let mut o = vec![0.0f32; nh * dv];
338    phase_step(
339        &thq,
340        &thk,
341        &v,
342        &w.decay,
343        kap.as_deref(),
344        cfg,
345        state,
346        &mut o,
347        w.phase_delta,
348    );
349
350    let mut out = vec![0.0f32; cfg.hidden_size];
351    w.out_proj.matvec(&o, &mut out, pool);
352    out
353}
354
355/// Fused two-position forward (speculative verify). Lane 1 commits into
356/// `state` (its token is always committed); lane 2's tentative state
357/// goes into `scratch` — the caller swaps it in on draft acceptance and
358/// simply drops it on rejection.
359#[allow(clippy::too_many_arguments)]
360pub fn vmf_phase_pair(
361    x1: &[f32],
362    x2: &[f32],
363    w: &VmfPhaseWeights,
364    cfg: &VmfPhaseCfg,
365    state: &mut Vec<f32>,
366    scratch: &mut Vec<f32>,
367    pool: Option<&Pool>,
368) -> (Vec<f32>, Vec<f32>) {
369    if w.conv.is_none() && state.len() != cfg.state_len() {
370        *state = vec![0f32; cfg.state_len()];
371    }
372    // Lane 1 commits its ring advance into the real state; lane 2 works on
373    // the tentative copy exactly like the phase state itself.
374    let xc1 = conv_in(x1, w, cfg, state);
375    let x1 = xc1.as_deref().unwrap_or(x1);
376    let (xc2, x2raw) = if w.conv.is_some() {
377        let mut tmp = state.clone();
378        (conv_in(x2, w, cfg, &mut tmp), Some(x2))
379    } else {
380        (None, None)
381    };
382    let x2 = xc2.as_deref().unwrap_or(x2);
383    let (nh, nph, dv) = (cfg.num_heads, cfg.nphase, cfg.value_head_dim);
384
385    let mut thq1 = vec![0.0f32; nh * nph];
386    let mut thq2 = vec![0.0f32; nh * nph];
387    w.thq.matvec2(x1, x2, &mut thq1, &mut thq2, pool);
388    let mut thk1 = vec![0.0f32; nh * nph];
389    let mut thk2 = vec![0.0f32; nh * nph];
390    w.thk.matvec2(x1, x2, &mut thk1, &mut thk2, pool);
391    let mut v1 = vec![0.0f32; nh * dv];
392    let mut v2 = vec![0.0f32; nh * dv];
393    w.v_proj.matvec2(x1, x2, &mut v1, &mut v2, pool);
394
395    // Lane 1 commits into the real state.
396    let kap1 = kappa_of(x1, w, nh, pool);
397    let mut o1 = vec![0.0f32; nh * dv];
398    phase_step(
399        &thq1,
400        &thk1,
401        &v1,
402        &w.decay,
403        kap1.as_deref(),
404        cfg,
405        state,
406        &mut o1,
407        w.phase_delta,
408    );
409
410    // Lane 2 runs on a copy — tentative until the draft is verified.
411    let kap2 = kappa_of(x2, w, nh, pool);
412    scratch.clear();
413    scratch.extend_from_slice(state);
414    let mut o2 = vec![0.0f32; nh * dv];
415    phase_step(
416        &thq2,
417        &thk2,
418        &v2,
419        &w.decay,
420        kap2.as_deref(),
421        cfg,
422        scratch,
423        &mut o2,
424        w.phase_delta,
425    );
426    // The scratch copy above re-took state's ring (advanced only through
427    // x1); commit x2's advance so an accepted draft leaves a correct ring.
428    if let Some(xr) = x2raw {
429        conv_ring_push(xr, cfg.hidden_size, cfg.state_len(), scratch);
430    }
431
432    let mut out1 = vec![0.0f32; cfg.hidden_size];
433    let mut out2 = vec![0.0f32; cfg.hidden_size];
434    w.out_proj.matvec2(&o1, &o2, &mut out1, &mut out2, pool);
435    (out1, out2)
436}
437
438// ───────────────────────── GatedDeltaNet (faithful vendor operator) ─────────────────────────
439
440/// Weights of one GatedDeltaNet layer (`model.layers.{i}.linear_attn.*`,
441/// names 1:1 with the source model — no fold, no training).
442pub struct GdnWeights {
443    /// [2·nk·dk + nv·dv, hidden] — fused q/k/v projection
444    pub in_proj_qkv: QTensor,
445    /// [nv·dv, hidden] — output-gate projection z
446    pub in_proj_z: QTensor,
447    /// [nv, hidden] — decay modulation a
448    pub in_proj_a: QTensor,
449    /// [nv, hidden] — write-strength b (β = σ(b))
450    pub in_proj_b: QTensor,
451    /// [c_dim · kk] — depthwise causal conv taps, flattened [c][tap]
452    pub conv1d: Vec<f32>,
453    /// [nv]
454    pub a_log: Vec<f32>,
455    /// [nv]
456    pub dt_bias: Vec<f32>,
457    /// [dv] — gated RMSNorm weight (plain x̂·w, validated by the oracle)
458    pub norm: Vec<f32>,
459    /// [hidden, nv·dv]
460    pub out_proj: QTensor,
461}
462
463#[derive(Clone, Copy)]
464pub struct GdnCfg {
465    pub num_v_heads: usize,
466    pub num_k_heads: usize,
467    pub key_head_dim: usize,
468    pub value_head_dim: usize,
469    pub conv_kernel: usize,
470    pub hidden_size: usize,
471    pub rms_eps: f64,
472    /// Qwen4-exp trains the output gate as sigmoid(z); older GDN families
473    /// use SiLU(z). This is part of the operator, not a sampling option.
474    pub output_gate_sigmoid: bool,
475}
476
477impl GdnCfg {
478    pub fn conv_dim(&self) -> usize {
479        2 * self.num_k_heads * self.key_head_dim + self.num_v_heads * self.value_head_dim
480    }
481
482    /// Packed state: [conv ring (kk−1)·c_dim | S nv·dk·dv], one Vec<f64>
483    /// so the speculative scratch-swap moves ring and recurrent state together.
484    pub fn state_len(&self) -> usize {
485        (self.conv_kernel - 1) * self.conv_dim()
486            + self.num_v_heads * self.key_head_dim * self.value_head_dim
487    }
488}
489
490fn softplus(x: f64) -> f64 {
491    if x > 20.0 { x } else { x.exp().ln_1p() }
492}
493
494fn sigmoid(x: f64) -> f64 {
495    1.0 / (1.0 + (-x).exp())
496}
497
498fn silu(x: f64) -> f64 {
499    x / (1.0 + (-x).exp())
500}
501
502/// `*mut f32` that may cross worker threads; safety comes from the
503/// disjoint (head, element) ranges each worker writes.
504#[derive(Clone, Copy)]
505struct SendMutF32(*mut f32);
506unsafe impl Send for SendMutF32 {}
507unsafe impl Sync for SendMutF32 {}
508
509/// One recurrent step given the raw (pre-conv) projections of this
510/// position. Advances the packed state (conv ring + S) and writes the
511/// gated per-head output into `of` [nv·dv].
512///
513/// The recurrent-state math runs in f32 (the vendor operator's own dtype —
514/// `mamba_ssm_dtype: float32` in the source configs; the old f64 was
515/// over-precision at 4× the traffic and no SIMD). The two S passes are
516/// element-wise in `dj` with no cross-lane reduction, so LLVM
517/// auto-vectorizes them (fmla on NEON, FMA on AVX2). Heads are
518/// independent given the conv output and run across the pool — on a
519/// Qwen3.5-27B this loop is 48 heads × 128×128 × 48 layers per token,
520/// the single biggest serial block in the hybrid's decode.
521#[allow(clippy::too_many_arguments)]
522fn gdn_step(
523    qkv: &[f32],
524    z: &[f32],
525    a: &[f32],
526    b: &[f32],
527    w: &GdnWeights,
528    cfg: &GdnCfg,
529    state: &mut [f32],
530    of: &mut [f32],
531    pool: Option<&Pool>,
532) {
533    let perf_t0 = perf_enabled().then(std::time::Instant::now);
534    let (nv, nk, dk, dv, kk) = (
535        cfg.num_v_heads,
536        cfg.num_k_heads,
537        cfg.key_head_dim,
538        cfg.value_head_dim,
539        cfg.conv_kernel,
540    );
541    let c_dim = cfg.conv_dim();
542    let (kd, rep) = (nk * dk, nv / nk);
543    let (ring, s_all) = state.split_at_mut((kk - 1) * c_dim);
544
545    // Depthwise causal conv over [ring…, current] + SiLU. Taps are
546    // ordered oldest→newest; tap kk−1 multiplies the current position.
547    // (Tiny: c_dim × kk — f64 accumulation kept.)
548    let mut cq = vec![0f32; c_dim];
549    for c in 0..c_dim {
550        let taps = &w.conv1d[c * kk..(c + 1) * kk];
551        let mut acc = qkv[c] as f64 * taps[kk - 1] as f64;
552        for j in 0..kk - 1 {
553            acc += ring[j * c_dim + c] as f64 * taps[j] as f64;
554        }
555        cq[c] = silu(acc) as f32;
556    }
557    // Ring shift: drop the oldest position, append the raw current one.
558    if kk > 1 {
559        ring.copy_within(c_dim.., 0);
560        let tail = (kk - 2) * c_dim;
561        ring[tail..tail + c_dim].copy_from_slice(&qkv[..c_dim]);
562    }
563
564    let cq = &cq;
565    let s_ptr = SendMutF32(s_all.as_mut_ptr());
566    let of_ptr = SendMutF32(of.as_mut_ptr());
567    let head_range = |h0: usize, h1: usize| {
568        // Rebind the Sync wrappers whole — edition-2021 disjoint capture
569        // would otherwise grab the raw `.0` fields and lose Send/Sync.
570        let (s_ptr, of_ptr) = (s_ptr, of_ptr);
571        // Per-worker scratch, recycled across calls (thread-local freelists).
572        let mut kv = crate::attention::take_buf(dv);
573        let mut delta = crate::attention::take_buf(dv);
574        let mut o = crate::attention::take_buf(dv);
575        let mut kf = crate::attention::take_buf(dk);
576        let mut qf = crate::attention::take_buf(dk);
577        for h in h0..h1 {
578            let ko = h / rep; // source q/k head (GQA)
579            let (qs, ks) = (ko * dk, kd + ko * dk);
580            // l2-normalize q and k; q additionally scaled by 1/√dk.
581            let (mut nq, mut nkn) = (0f64, 0f64);
582            for d in 0..dk {
583                nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
584                nkn += (cq[ks + d] as f64) * (cq[ks + d] as f64);
585            }
586            let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
587            let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
588            for d in 0..dk {
589                qf[d] = cq[qs + d] * invq;
590                kf[d] = cq[ks + d] * invk;
591            }
592
593            let g = (-(w.a_log[h] as f64).exp() * softplus(a[h] as f64 + w.dt_bias[h] as f64)).exp()
594                as f32;
595            let beta = sigmoid(b[h] as f64) as f32;
596
597            // SAFETY: disjoint per-head S and output slices per worker.
598            let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
599            let oh = unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
600            let vt = &cq[2 * kd + h * dv..2 * kd + (h + 1) * dv];
601
602            // S ← g·S;  kv = kᵀS;  S += k ⊗ β(v − kv);  o = qᵀS —
603            // algebraically regrouped so S is READ twice and WRITTEN
604            // once: kv over S_old (then ×g), one fused update+query pass.
605            kv[..dv].fill(0.0);
606            for di in 0..dk {
607                let kfd = kf[di];
608                let row = &s[di * dv..(di + 1) * dv];
609                for dj in 0..dv {
610                    kv[dj] += row[dj] * kfd; // elementwise in dj → SIMD
611                }
612            }
613            for dj in 0..dv {
614                delta[dj] = (vt[dj] - g * kv[dj]) * beta;
615            }
616            o[..dv].fill(0.0);
617            for di in 0..dk {
618                let kfd = kf[di];
619                let qfd = qf[di];
620                let row = &mut s[di * dv..(di + 1) * dv];
621                for dj in 0..dv {
622                    let cell = g * row[dj] + kfd * delta[dj];
623                    row[dj] = cell;
624                    o[dj] += qfd * cell; // elementwise in dj → SIMD
625                }
626            }
627            // Gated RMSNorm per head: x̂·w·silu(z) (oracle-validated form).
628            let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
629            let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
630            for dj in 0..dv {
631                let gate = if cfg.output_gate_sigmoid {
632                    sigmoid(z[h * dv + dj] as f64)
633                } else {
634                    silu(z[h * dv + dj] as f64)
635                };
636                oh[dj] = ((o[dj] as f64 * inv) * w.norm[dj] as f64 * gate) as f32;
637            }
638        }
639        crate::attention::recycle_buf(&mut kv);
640        crate::attention::recycle_buf(&mut delta);
641        crate::attention::recycle_buf(&mut o);
642        crate::attention::recycle_buf(&mut kf);
643        crate::attention::recycle_buf(&mut qf);
644    };
645    match pool {
646        Some(pool) if nv >= 4 => pool.run(&|widx, n| {
647            let chunk = nv.div_ceil(n);
648            let h0 = (widx * chunk).min(nv);
649            let h1 = (h0 + chunk).min(nv);
650            if h0 < h1 {
651                head_range(h0, h1);
652            }
653        }),
654        _ => head_range(0, nv),
655    }
656    if let Some(t0) = perf_t0 {
657        PERF_GDN_STEP_CALLS.fetch_add(1, Ordering::Relaxed);
658        PERF_GDN_STEP_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
659    }
660}
661
662/// Forward one position through a GatedDeltaNet layer, advancing `state`.
663pub fn gdn_forward(
664    x: &[f32],
665    w: &GdnWeights,
666    cfg: &GdnCfg,
667    state: &mut Vec<f32>,
668    pool: Option<&Pool>,
669) -> Vec<f32> {
670    let perf_t0 = perf_enabled().then(std::time::Instant::now);
671    if state.len() != cfg.state_len() {
672        *state = vec![0f32; cfg.state_len()];
673    }
674    let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
675
676    let mut qkv = vec![0.0f32; c_dim];
677    let mut z = vec![0.0f32; vd];
678    let mut a = vec![0.0f32; cfg.num_v_heads];
679    let mut b = vec![0.0f32; cfg.num_v_heads];
680    // D5: two heavy projections (the GDN mixer is ~half a hybrid layer's
681    // bytes) — one GPU submission; a/b are tiny and stay on CPU. The
682    // Batch probe arbitrates GPU vs the fused-CPU dispatch per machine.
683    let cpu_projs = |qkv: &mut Vec<f32>, z: &mut Vec<f32>, a: &mut Vec<f32>, b: &mut Vec<f32>| {
684        QTensor::matvec_many(
685            [&w.in_proj_qkv, &w.in_proj_z, &w.in_proj_a, &w.in_proj_b],
686            x,
687            [
688                qkv.as_mut_slice(),
689                z.as_mut_slice(),
690                a.as_mut_slice(),
691                b.as_mut_slice(),
692            ],
693            pool,
694        );
695    };
696    let mut done = false;
697    if crate::gpu::enabled_here() && gdn_projs_eligible(w) {
698        match crate::gpu::probe_arm(crate::gpu::OpClass::Batch) {
699            crate::gpu::ProbeArm::Gpu => {
700                let t0 = std::time::Instant::now();
701                if gdn_projs_gpu(w, x, &mut qkv, &mut z) {
702                    crate::gpu::probe_record(crate::gpu::OpClass::Batch, true, t0.elapsed());
703                    w.in_proj_a.matvec(x, &mut a, pool);
704                    w.in_proj_b.matvec(x, &mut b, pool);
705                    done = true;
706                } else {
707                    crate::gpu::probe_note_decline(crate::gpu::OpClass::Batch);
708                }
709            }
710            crate::gpu::ProbeArm::CpuTimed => {
711                let t0 = std::time::Instant::now();
712                crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
713                crate::gpu::probe_record(crate::gpu::OpClass::Batch, false, t0.elapsed());
714                done = true;
715            }
716            crate::gpu::ProbeArm::Cpu => {
717                crate::gpu::cpu_scope(|| cpu_projs(&mut qkv, &mut z, &mut a, &mut b));
718                done = true;
719            }
720        }
721    }
722    if !done {
723        cpu_projs(&mut qkv, &mut z, &mut a, &mut b);
724    }
725
726    let mut of = vec![0.0f32; vd];
727    gdn_step(&qkv, &z, &a, &b, w, cfg, state, &mut of, pool);
728
729    let mut out = vec![0.0f32; cfg.hidden_size];
730    w.out_proj.matvec(&of, &mut out, pool);
731    if let Some(t0) = perf_t0 {
732        PERF_GDN_FORWARD_CALLS.fetch_add(1, Ordering::Relaxed);
733        PERF_GDN_FORWARD_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
734    }
735    out
736}
737
738/// Batched GDN forward (prefill-GEMM): the qkv/z/a/b and out_proj
739/// projections are matmat over the batch (a weight row once per chunk),
740/// the gdn_step recurrence runs sequentially over positions (state is the
741/// same as the sequential path; the math is elementwise identical).
742pub fn gdn_forward_batch(
743    xs: &[f32],
744    b: usize,
745    w: &GdnWeights,
746    cfg: &GdnCfg,
747    state: &mut Vec<f32>,
748    pool: Option<&Pool>,
749) -> Vec<f32> {
750    let perf_t0 = perf_enabled().then(std::time::Instant::now);
751    if state.len() != cfg.state_len() {
752        *state = vec![0f32; cfg.state_len()];
753    }
754    let (c_dim, vd) = (cfg.conv_dim(), cfg.num_v_heads * cfg.value_head_dim);
755    let nv = cfg.num_v_heads;
756
757    let mut qkv = vec![0.0f32; b * c_dim];
758    w.in_proj_qkv.matmat(xs, b, &mut qkv, pool);
759    let mut z = vec![0.0f32; b * vd];
760    w.in_proj_z.matmat(xs, b, &mut z, pool);
761    let mut a = vec![0.0f32; b * nv];
762    w.in_proj_a.matmat(xs, b, &mut a, pool);
763    let mut bb = vec![0.0f32; b * nv];
764    w.in_proj_b.matmat(xs, b, &mut bb, pool);
765
766    let mut of = vec![0.0f32; b * vd];
767    for bi in 0..b {
768        gdn_step(
769            &qkv[bi * c_dim..(bi + 1) * c_dim],
770            &z[bi * vd..(bi + 1) * vd],
771            &a[bi * nv..(bi + 1) * nv],
772            &bb[bi * nv..(bi + 1) * nv],
773            w,
774            cfg,
775            state,
776            &mut of[bi * vd..(bi + 1) * vd],
777            pool,
778        );
779    }
780    let mut out = vec![0.0f32; b * cfg.hidden_size];
781    w.out_proj.matmat(&of, b, &mut out, pool);
782    if std::env::var("CMF_GDN_TRACE").is_ok() {
783        let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
784        eprintln!(
785            "gdn-batch b={b}: |x0|={:.5} |qkv0|={:.5} |z0|={:.5} |a0|={:.5} |b0|={:.5} |of0|={:.5} |out0|={:.5} |state|={:.5}",
786            n(&xs[..cfg.hidden_size]),
787            n(&qkv[..c_dim]),
788            n(&z[..vd]),
789            n(&a[..nv]),
790            n(&bb[..nv]),
791            n(&of[..vd]),
792            n(&out[..cfg.hidden_size]),
793            n(state)
794        );
795    }
796    if let Some(t0) = perf_t0 {
797        PERF_GDN_BATCH_CALLS.fetch_add(1, Ordering::Relaxed);
798        PERF_GDN_BATCH_NS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
799    }
800    out
801}
802
803/// Aggregate host-side GDN costs for the bounded warm-decode profile.
804pub fn perf_report() {
805    if !perf_enabled() {
806        return;
807    }
808    let fc = PERF_GDN_FORWARD_CALLS.load(Ordering::Relaxed);
809    let bc = PERF_GDN_BATCH_CALLS.load(Ordering::Relaxed);
810    let sc = PERF_GDN_STEP_CALLS.load(Ordering::Relaxed);
811    eprintln!(
812        "[perf-gdn] forward_calls={} forward_ms={:.3} forward_ms_per_call={:.3} batch_calls={} batch_ms={:.3} step_calls={} step_ms={:.3} step_ms_per_call={:.3}",
813        fc,
814        PERF_GDN_FORWARD_NS.load(Ordering::Relaxed) as f64 / 1e6,
815        PERF_GDN_FORWARD_NS.load(Ordering::Relaxed) as f64 / 1e6 / fc.max(1) as f64,
816        bc,
817        PERF_GDN_BATCH_NS.load(Ordering::Relaxed) as f64 / 1e6,
818        sc,
819        PERF_GDN_STEP_NS.load(Ordering::Relaxed) as f64 / 1e6,
820        PERF_GDN_STEP_NS.load(Ordering::Relaxed) as f64 / 1e6 / sc.max(1) as f64,
821    );
822}
823
824/// GDN qkv+z GPU eligibility: q1 mixers offload by default (the CPU q1
825/// kernel is compute-bound); q8 stays opt-in via CMF_GPU_GDN=1 (measured
826/// neutral). The probe in `gdn_forward` still arbitrates either way.
827fn gdn_projs_eligible(w: &GdnWeights) -> bool {
828    // Any tile-embedded-scale layout, not just q1: the batched matvec now has
829    // kernels for all of them, and the probe arbitrates whether it pays.
830    // `CMF_GPU_GDN=0` forces the CPU projections back — the switch exists so
831    // the two arms can be compared inside ONE run: this machine throttles far
832    // enough that measurements minutes apart are not comparable.
833    if std::env::var("CMF_GPU_GDN")
834        .map(|v| v == "0")
835        .unwrap_or(false)
836    {
837        return false;
838    }
839    w.in_proj_qkv.is_q1()
840        || w.in_proj_qkv.q4t_parts().is_some()
841        || w.in_proj_qkv.q4tp_parts().is_some()
842        || std::env::var("CMF_GPU_GDN")
843            .map(|v| v == "1")
844            .unwrap_or(false)
845}
846
847/// GDN qkv+z on GPU in a single submission (independent matvecs of one input).
848fn gdn_projs_gpu(w: &GdnWeights, x: &[f32], qkv: &mut [f32], z: &mut [f32]) -> bool {
849    use crate::gpu::matvec_batch;
850    use crate::qtensor::QTensor;
851    if !crate::gpu::enabled_here() {
852        return false;
853    }
854    fn part<'a>(
855        t: &'a QTensor,
856        x: &[f32],
857    ) -> Option<(
858        std::sync::Arc<cortiq_core::CmfModel>,
859        crate::gpu::BatchJob<'a>,
860    )> {
861        use crate::gpu::BatchJob;
862        use crate::qtensor::prescale;
863        use cortiq_core::TensorDtype;
864        match t {
865            QTensor::Mapped {
866                model,
867                idx,
868                dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
869                rows,
870                cols,
871                row_scale,
872                col_field,
873                ..
874            } => Some((
875                model.clone(),
876                BatchJob {
877                    idx: *idx,
878                    rows: *rows,
879                    cols: *cols,
880                    row_scale,
881                    xs: prescale(x, col_field, *dt).into_owned(),
882                    layout: crate::gpu::BatchLayout::Q8,
883                },
884            )),
885            QTensor::Mapped {
886                model,
887                idx,
888                dtype: TensorDtype::Q1,
889                rows,
890                cols,
891                ..
892            } => Some((
893                model.clone(),
894                BatchJob {
895                    idx: *idx,
896                    rows: *rows,
897                    cols: *cols,
898                    row_scale: &[],
899                    xs: x.to_vec(),
900                    layout: crate::gpu::BatchLayout::Q1,
901                },
902            )),
903            // q4t/q4tp: same tile-embedded-scale contract as q1, different
904            // stride. Missing here, a hybrid model's GDN projections fell to
905            // the CPU on every layer while its experts rode the device.
906            QTensor::Mapped {
907                model,
908                idx,
909                dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
910                rows,
911                cols,
912                ..
913            } => Some((
914                model.clone(),
915                BatchJob {
916                    idx: *idx,
917                    rows: *rows,
918                    cols: *cols,
919                    row_scale: &[],
920                    xs: x.to_vec(),
921                    layout: if *dt == TensorDtype::Q4TiledP {
922                        crate::gpu::BatchLayout::Q4tp
923                    } else {
924                        crate::gpu::BatchLayout::Q4t
925                    },
926                },
927            )),
928            _ => None,
929        }
930    }
931    let Some((model, jq)) = part(&w.in_proj_qkv, x) else {
932        return false;
933    };
934    let Some((_, jz)) = part(&w.in_proj_z, x) else {
935        return false;
936    };
937    matvec_batch(&model, &[jq, jz], &mut [qkv, z])
938}
939
940/// Fused two-position forward (speculative verify): lane 1 commits into
941/// `state`, lane 2 is tentative in `scratch` (ring + S move together).
942#[allow(clippy::too_many_arguments)]
943pub fn gdn_pair(
944    x1: &[f32],
945    x2: &[f32],
946    w: &GdnWeights,
947    cfg: &GdnCfg,
948    state: &mut Vec<f32>,
949    scratch: &mut Vec<f32>,
950    pool: Option<&Pool>,
951) -> (Vec<f32>, Vec<f32>) {
952    if state.len() != cfg.state_len() {
953        *state = vec![0f32; cfg.state_len()];
954    }
955    let (c_dim, vd, nv) = (
956        cfg.conv_dim(),
957        cfg.num_v_heads * cfg.value_head_dim,
958        cfg.num_v_heads,
959    );
960
961    let mut qkv1 = vec![0.0f32; c_dim];
962    let mut qkv2 = vec![0.0f32; c_dim];
963    w.in_proj_qkv.matvec2(x1, x2, &mut qkv1, &mut qkv2, pool);
964    let mut z1 = vec![0.0f32; vd];
965    let mut z2 = vec![0.0f32; vd];
966    w.in_proj_z.matvec2(x1, x2, &mut z1, &mut z2, pool);
967    let mut a1 = vec![0.0f32; nv];
968    let mut a2 = vec![0.0f32; nv];
969    w.in_proj_a.matvec2(x1, x2, &mut a1, &mut a2, pool);
970    let mut b1 = vec![0.0f32; nv];
971    let mut b2 = vec![0.0f32; nv];
972    w.in_proj_b.matvec2(x1, x2, &mut b1, &mut b2, pool);
973
974    let mut of1 = vec![0.0f32; vd];
975    gdn_step(&qkv1, &z1, &a1, &b1, w, cfg, state, &mut of1, pool);
976
977    scratch.clear();
978    scratch.extend_from_slice(state);
979    let mut of2 = vec![0.0f32; vd];
980    gdn_step(&qkv2, &z2, &a2, &b2, w, cfg, scratch, &mut of2, pool);
981
982    let mut out1 = vec![0.0f32; cfg.hidden_size];
983    let mut out2 = vec![0.0f32; cfg.hidden_size];
984    w.out_proj.matvec2(&of1, &of2, &mut out1, &mut out2, pool);
985    (out1, out2)
986}
987
988// ───────────────────────── ShortConv (LFM2 gated short convolution) ─────────────────────────
989
990/// Weights of one LFM2 short-convolution mixer
991/// (`model.layers.{i}.short_conv.*`, renamed from the vendor `conv.*` at
992/// convert time). No recurrent mixer state — the only state is the causal
993/// conv ring (the last `kernel−1` gated inputs per channel).
994pub struct ShortConvWeights {
995    /// [3·hidden, hidden] — fused (B, C, x) projection.
996    pub in_proj: QTensor,
997    /// [hidden · kernel] depthwise conv taps, flattened `[channel][tap]`
998    /// (the source `[hidden, 1, kernel]` with the singleton group axis
999    /// dropped). Tap `kernel−1` multiplies the current position.
1000    pub conv: Vec<f32>,
1001    /// [hidden, hidden] — output projection.
1002    pub out_proj: QTensor,
1003}
1004
1005#[derive(Clone, Copy)]
1006pub struct ShortConvCfg {
1007    pub hidden_size: usize,
1008    /// Conv kernel width `L` (`conv_L_cache`; LFM2 uses 3).
1009    pub kernel: usize,
1010}
1011
1012impl ShortConvCfg {
1013    /// Conv ring: the last `kernel−1` gated inputs per channel.
1014    pub fn state_len(&self) -> usize {
1015        (self.kernel - 1) * self.hidden_size
1016    }
1017}
1018
1019/// One position through the gated conv, given the fused projection
1020/// `bcx = in_proj·x` [3·hidden] = [B | C | x]. Advances the conv ring and
1021/// writes the gated conv output `y = C ⊙ conv(B ⊙ x)` [hidden] into `y`.
1022///
1023/// The conv is PyTorch's causal depthwise `Conv1d(padding=kernel−1)`
1024/// truncated to the current length: for tap `k`, weight `w[c][k]` pairs
1025/// with the input `kernel−1−k` steps in the past, so `w[c][kernel−1]` is
1026/// the current position. The ring holds `in[t−1] … in[t−(kernel−1)]` at
1027/// slots `0 … kernel−2`.
1028fn short_conv_step(
1029    bcx: &[f32],
1030    conv: &[f32],
1031    cfg: &ShortConvCfg,
1032    ring_state: &mut [f32],
1033    y: &mut [f32],
1034) {
1035    let (h, k) = (cfg.hidden_size, cfg.kernel);
1036    let ring = k - 1;
1037    let (bg, cg, xg) = (&bcx[0..h], &bcx[h..2 * h], &bcx[2 * h..3 * h]);
1038    for c in 0..h {
1039        let bx = bg[c] * xg[c];
1040        let wc = &conv[c * k..(c + 1) * k];
1041        // Current tap, then the past taps read from the channel's ring.
1042        let mut acc = wc[k - 1] * bx;
1043        let rc = &mut ring_state[c * ring..c * ring + ring];
1044        for s in 0..ring {
1045            acc += wc[k - 2 - s] * rc[s];
1046        }
1047        y[c] = cg[c] * acc;
1048        // Shift newest-in-front: slot 0 becomes the just-seen input.
1049        for s in (1..ring).rev() {
1050            rc[s] = rc[s - 1];
1051        }
1052        if ring > 0 {
1053            rc[0] = bx;
1054        }
1055    }
1056}
1057
1058/// Forward one position through a short-conv layer, advancing `state`.
1059pub fn short_conv_forward(
1060    x: &[f32],
1061    w: &ShortConvWeights,
1062    cfg: &ShortConvCfg,
1063    state: &mut Vec<f32>,
1064    pool: Option<&Pool>,
1065) -> Vec<f32> {
1066    if state.len() != cfg.state_len() {
1067        *state = vec![0f32; cfg.state_len()];
1068    }
1069    let h = cfg.hidden_size;
1070    let mut bcx = vec![0.0f32; 3 * h];
1071    w.in_proj.matvec(x, &mut bcx, pool);
1072    let mut y = vec![0.0f32; h];
1073    short_conv_step(&bcx, &w.conv, cfg, state, &mut y);
1074    let mut out = vec![0.0f32; h];
1075    w.out_proj.matvec(&y, &mut out, pool);
1076    out
1077}
1078
1079/// Batched short-conv forward (prefill-GEMM): in_proj/out_proj are matmat
1080/// over the chunk (a weight row streamed once), the conv walks the
1081/// positions in order — the chunk is contiguous, so the ring state is
1082/// exactly the sequential path's and the math is elementwise identical.
1083pub fn short_conv_forward_batch(
1084    xs: &[f32],
1085    b: usize,
1086    w: &ShortConvWeights,
1087    cfg: &ShortConvCfg,
1088    state: &mut Vec<f32>,
1089    pool: Option<&Pool>,
1090) -> Vec<f32> {
1091    if state.len() != cfg.state_len() {
1092        *state = vec![0f32; cfg.state_len()];
1093    }
1094    let h = cfg.hidden_size;
1095    let mut bcx = vec![0.0f32; b * 3 * h];
1096    w.in_proj.matmat(xs, b, &mut bcx, pool);
1097    let mut y = vec![0.0f32; b * h];
1098    for bi in 0..b {
1099        short_conv_step(
1100            &bcx[bi * 3 * h..(bi + 1) * 3 * h],
1101            &w.conv,
1102            cfg,
1103            state,
1104            &mut y[bi * h..(bi + 1) * h],
1105        );
1106    }
1107    let mut out = vec![0.0f32; b * h];
1108    w.out_proj.matmat(&y, b, &mut out, pool);
1109    out
1110}
1111
1112/// Fused two-position forward (speculative verify). Lane 1 commits into
1113/// `state`; lane 2's tentative ring goes into `scratch` — swapped in on
1114/// draft acceptance, dropped on rejection. LFM2 ships no MTP head, so this
1115/// is exercised only by the pair-fusion micro-benchmark; kept correct.
1116#[allow(clippy::too_many_arguments)]
1117pub fn short_conv_pair(
1118    x1: &[f32],
1119    x2: &[f32],
1120    w: &ShortConvWeights,
1121    cfg: &ShortConvCfg,
1122    state: &mut Vec<f32>,
1123    scratch: &mut Vec<f32>,
1124    pool: Option<&Pool>,
1125) -> (Vec<f32>, Vec<f32>) {
1126    if state.len() != cfg.state_len() {
1127        *state = vec![0f32; cfg.state_len()];
1128    }
1129    let h = cfg.hidden_size;
1130    let mut bcx1 = vec![0.0f32; 3 * h];
1131    let mut bcx2 = vec![0.0f32; 3 * h];
1132    w.in_proj.matvec2(x1, x2, &mut bcx1, &mut bcx2, pool);
1133
1134    let mut y1 = vec![0.0f32; h];
1135    short_conv_step(&bcx1, &w.conv, cfg, state, &mut y1);
1136    scratch.clear();
1137    scratch.extend_from_slice(state);
1138    let mut y2 = vec![0.0f32; h];
1139    short_conv_step(&bcx2, &w.conv, cfg, scratch, &mut y2);
1140
1141    let mut out1 = vec![0.0f32; h];
1142    let mut out2 = vec![0.0f32; h];
1143    w.out_proj.matvec2(&y1, &y2, &mut out1, &mut out2, pool);
1144    (out1, out2)
1145}
1146
1147// ─── Kimi Delta Attention (KDA) ─────────────────────────────────────────
1148//
1149// Kimi Linear / Kimi-K3 linear mixer (reference: FLA naive_recurrent_kda
1150// + moonshotai modeling_kimi.py). Differences from GatedDeltaNet above:
1151// separate q/k/v projections each behind its OWN causal depthwise short
1152// convolution; the delta-rule decay is a PER-CHANNEL vector (diagonal)
1153// instead of a per-head scalar; the decay pre-activation comes from a
1154// low-rank projection f_b(f_a(x)); and the output gate norm uses
1155// sigmoid, not SiLU.
1156
1157pub struct KdaWeights {
1158    /// [nh·dk, hidden]
1159    pub q_proj: QTensor,
1160    /// [nh·dk, hidden]
1161    pub k_proj: QTensor,
1162    /// [nh·dv, hidden]
1163    pub v_proj: QTensor,
1164    /// [nh·dk × kk] — depthwise taps, oldest→newest (see GdnWeights.conv1d)
1165    pub conv_q: Vec<f32>,
1166    pub conv_k: Vec<f32>,
1167    /// [nh·dv × kk]
1168    pub conv_v: Vec<f32>,
1169    /// [rank, hidden] — low-rank decay projection, stage 1
1170    pub f_a: QTensor,
1171    /// [nh·dk, rank] — stage 2
1172    pub f_b: QTensor,
1173    /// [nh·dk]
1174    pub dt_bias: Vec<f32>,
1175    /// [nh] per-head (Kimi-Linear-48B) | [dk] per-dim (Kimi-K3) |
1176    /// [nh·dk] full — broadcast resolved by length.
1177    pub a_log: Vec<f32>,
1178    /// [nh, hidden] — β = σ(b_proj·x) per head
1179    pub b_proj: QTensor,
1180    /// Output gate: full-rank g_proj (K3) or low-rank g_b(g_a(x)) (48B).
1181    pub gate: KdaOutGate,
1182    /// [dv] — gated RMSNorm weight (per head over head_v_dim)
1183    pub o_norm: Vec<f32>,
1184    /// [hidden, nh·dv]
1185    pub o_proj: QTensor,
1186    /// Some(lb): log-decay = lb·σ(exp(A)·(f+bias)) (K3, lb=−5);
1187    /// None: −exp(A)·softplus(f+bias) (Kimi-Linear-48B).
1188    pub gate_lower_bound: Option<f32>,
1189}
1190
1191pub enum KdaOutGate {
1192    /// [nh·dv, hidden]
1193    Full(QTensor),
1194    /// g_a [rank, hidden], g_b [nh·dv, rank]
1195    LowRank(QTensor, QTensor),
1196}
1197
1198#[derive(Clone, Copy)]
1199pub struct KdaCfg {
1200    pub num_heads: usize,
1201    pub head_k_dim: usize,
1202    pub head_v_dim: usize,
1203    pub conv_kernel: usize,
1204    pub hidden_size: usize,
1205    pub rms_eps: f64,
1206}
1207
1208impl KdaCfg {
1209    /// Packed state: [q ring | k ring | v ring | S nh·dk·dv], one Vec —
1210    /// same single-buffer convention as GdnCfg::state_len.
1211    pub fn state_len(&self) -> usize {
1212        let (nh, dk, dv, kk) = (
1213            self.num_heads,
1214            self.head_k_dim,
1215            self.head_v_dim,
1216            self.conv_kernel,
1217        );
1218        (kk - 1) * (2 * nh * dk + nh * dv) + nh * dk * dv
1219    }
1220}
1221
1222/// Depthwise causal conv over [ring…, current] + SiLU, then ring shift.
1223/// Taps oldest→newest, tap kk−1 multiplies the current position.
1224fn kda_conv(raw: &[f32], taps: &[f32], ring: &mut [f32], kk: usize, out: &mut [f32]) {
1225    let c_dim = raw.len();
1226    for c in 0..c_dim {
1227        let t = &taps[c * kk..(c + 1) * kk];
1228        let mut acc = raw[c] as f64 * t[kk - 1] as f64;
1229        for j in 0..kk - 1 {
1230            acc += ring[j * c_dim + c] as f64 * t[j] as f64;
1231        }
1232        out[c] = silu(acc) as f32;
1233    }
1234    if kk > 1 {
1235        ring.copy_within(c_dim.., 0);
1236        let tail = (kk - 2) * c_dim;
1237        ring[tail..tail + c_dim].copy_from_slice(raw);
1238    }
1239}
1240
1241/// Per-channel log-decay for head-channel (h, d): resolves the A_log
1242/// broadcast by length and applies the configured gate formula.
1243#[inline]
1244fn kda_log_decay(w: &KdaWeights, cfg: &KdaCfg, h: usize, d: usize, f: f32) -> f64 {
1245    let (nh, dk) = (cfg.num_heads, cfg.head_k_dim);
1246    let a = if w.a_log.len() == nh {
1247        w.a_log[h] as f64
1248    } else if w.a_log.len() == dk {
1249        w.a_log[d] as f64
1250    } else {
1251        w.a_log[h * dk + d] as f64
1252    };
1253    let raw = f as f64 + w.dt_bias[h * dk + d] as f64;
1254    match w.gate_lower_bound {
1255        Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
1256        None => -a.exp() * softplus(raw),
1257    }
1258}
1259
1260/// One recurrent step given this position's raw (pre-conv) projections.
1261/// Advances the packed state and writes the gated per-head output into
1262/// `of` [nh·dv]. Recurrence (FLA naive_recurrent_kda):
1263///   S ← Diag(exp(g))·S;  S += β·k ⊗ (v − kᵀS);  o = qᵀS
1264/// with q,k L2-normalized per head and q additionally scaled by 1/√dk —
1265/// regrouped into two S passes like gdn_step (per-channel decay folds
1266/// into the k readout of the first pass).
1267#[allow(clippy::too_many_arguments)]
1268fn kda_step(
1269    xq: &[f32],
1270    xk: &[f32],
1271    xv: &[f32],
1272    f: &[f32],
1273    b: &[f32],
1274    gate_out: &[f32],
1275    w: &KdaWeights,
1276    cfg: &KdaCfg,
1277    state: &mut [f32],
1278    of: &mut [f32],
1279    pool: Option<&Pool>,
1280) {
1281    let (nh, dk, dv, kk) = (
1282        cfg.num_heads,
1283        cfg.head_k_dim,
1284        cfg.head_v_dim,
1285        cfg.conv_kernel,
1286    );
1287    let (kd, vd) = (nh * dk, nh * dv);
1288    let ring_q_len = (kk - 1) * kd;
1289    let ring_v_len = (kk - 1) * vd;
1290    let (ring_q, rest) = state.split_at_mut(ring_q_len);
1291    let (ring_k, rest) = rest.split_at_mut(ring_q_len);
1292    let (ring_v, s_all) = rest.split_at_mut(ring_v_len);
1293
1294    let mut cq = vec![0f32; kd];
1295    let mut ck = vec![0f32; kd];
1296    let mut cv = vec![0f32; vd];
1297    kda_conv(xq, &w.conv_q, ring_q, kk, &mut cq);
1298    kda_conv(xk, &w.conv_k, ring_k, kk, &mut ck);
1299    kda_conv(xv, &w.conv_v, ring_v, kk, &mut cv);
1300
1301    let (cq, ck, cv) = (&cq, &ck, &cv);
1302    let s_ptr = SendMutF32(s_all.as_mut_ptr());
1303    let of_ptr = SendMutF32(of.as_mut_ptr());
1304    let head_range = |h0: usize, h1: usize| {
1305        let (s_ptr, of_ptr) = (s_ptr, of_ptr);
1306        let mut kv = crate::attention::take_buf(dv);
1307        let mut delta = crate::attention::take_buf(dv);
1308        let mut o = crate::attention::take_buf(dv);
1309        let mut kf = crate::attention::take_buf(dk);
1310        let mut qf = crate::attention::take_buf(dk);
1311        let mut gd = crate::attention::take_buf(dk);
1312        for h in h0..h1 {
1313            let qs = h * dk;
1314            // l2-normalize q and k; q additionally scaled by 1/√dk.
1315            let (mut nq, mut nkn) = (0f64, 0f64);
1316            for d in 0..dk {
1317                nq += (cq[qs + d] as f64) * (cq[qs + d] as f64);
1318                nkn += (ck[qs + d] as f64) * (ck[qs + d] as f64);
1319            }
1320            let invq = (1.0 / ((nq + 1e-6).sqrt() * (dk as f64).sqrt())) as f32;
1321            let invk = (1.0 / (nkn + 1e-6).sqrt()) as f32;
1322            for d in 0..dk {
1323                qf[d] = cq[qs + d] * invq;
1324                kf[d] = ck[qs + d] * invk;
1325                gd[d] = kda_log_decay(w, cfg, h, d, f[qs + d]).exp() as f32;
1326            }
1327            let beta = sigmoid(b[h] as f64) as f32;
1328
1329            // SAFETY: disjoint per-head S and output slices per worker.
1330            let s = unsafe { std::slice::from_raw_parts_mut(s_ptr.0.add(h * dk * dv), dk * dv) };
1331            let oh = unsafe { std::slice::from_raw_parts_mut(of_ptr.0.add(h * dv), dv) };
1332            let vt = &cv[h * dv..(h + 1) * dv];
1333
1334            // Pass 1: kv = kᵀ(Diag(gd)·S_old) — decay folded into k.
1335            kv[..dv].fill(0.0);
1336            for di in 0..dk {
1337                let kg = kf[di] * gd[di];
1338                let row = &s[di * dv..(di + 1) * dv];
1339                for dj in 0..dv {
1340                    kv[dj] += row[dj] * kg;
1341                }
1342            }
1343            for dj in 0..dv {
1344                delta[dj] = (vt[dj] - kv[dj]) * beta;
1345            }
1346            // Pass 2: S[di,:] = gd[di]·row + k[di]·delta;  o += q[di]·row.
1347            o[..dv].fill(0.0);
1348            for di in 0..dk {
1349                let (kfd, qfd, gdd) = (kf[di], qf[di], gd[di]);
1350                let row = &mut s[di * dv..(di + 1) * dv];
1351                for dj in 0..dv {
1352                    let cell = gdd * row[dj] + kfd * delta[dj];
1353                    row[dj] = cell;
1354                    o[dj] += qfd * cell;
1355                }
1356            }
1357            // Gated RMSNorm per head: x̂·w·σ(gate) — sigmoid, not SiLU.
1358            let ss: f64 = o[..dv].iter().map(|&v| (v as f64) * (v as f64)).sum();
1359            let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
1360            for dj in 0..dv {
1361                oh[dj] = ((o[dj] as f64 * inv)
1362                    * w.o_norm[dj] as f64
1363                    * sigmoid(gate_out[h * dv + dj] as f64)) as f32;
1364            }
1365        }
1366        crate::attention::recycle_buf(&mut kv);
1367        crate::attention::recycle_buf(&mut delta);
1368        crate::attention::recycle_buf(&mut o);
1369        crate::attention::recycle_buf(&mut kf);
1370        crate::attention::recycle_buf(&mut qf);
1371        crate::attention::recycle_buf(&mut gd);
1372    };
1373    match pool {
1374        Some(pool) if nh >= 4 => pool.run(&|widx, n| {
1375            let chunk = nh.div_ceil(n);
1376            let h0 = (widx * chunk).min(nh);
1377            let h1 = (h0 + chunk).min(nh);
1378            if h0 < h1 {
1379                head_range(h0, h1);
1380            }
1381        }),
1382        _ => head_range(0, nh),
1383    }
1384}
1385
1386/// Project one position's raw q/k/v/f/β/gate inputs (shared by the
1387/// single and batched forwards; `bi` selects the row when batched).
1388fn kda_gate_out(w: &KdaWeights, x: &[f32], vd: usize, pool: Option<&Pool>) -> Vec<f32> {
1389    let mut g = vec![0.0f32; vd];
1390    match &w.gate {
1391        KdaOutGate::Full(gp) => gp.matvec(x, &mut g, pool),
1392        KdaOutGate::LowRank(ga, gb) => {
1393            let mut low = vec![0.0f32; ga.rows()];
1394            ga.matvec(x, &mut low, pool);
1395            gb.matvec(&low, &mut g, pool);
1396        }
1397    }
1398    g
1399}
1400
1401/// Forward one position through a KDA layer, advancing `state`.
1402pub fn kda_forward(
1403    x: &[f32],
1404    w: &KdaWeights,
1405    cfg: &KdaCfg,
1406    state: &mut Vec<f32>,
1407    pool: Option<&Pool>,
1408) -> Vec<f32> {
1409    if state.len() != cfg.state_len() {
1410        *state = vec![0f32; cfg.state_len()];
1411    }
1412    let (nh, dk, dv) = (cfg.num_heads, cfg.head_k_dim, cfg.head_v_dim);
1413    let (kd, vd) = (nh * dk, nh * dv);
1414
1415    let mut xq = vec![0.0f32; kd];
1416    let mut xk = vec![0.0f32; kd];
1417    let mut xv = vec![0.0f32; vd];
1418    let mut fl = vec![0.0f32; w.f_a.rows()];
1419    let mut b = vec![0.0f32; nh];
1420    // KDA's beta projection uses the same hidden input as q/k/v/f_a.  When
1421    // it shares their Q4TP layout, include it in the virtual row space so
1422    // the projection group pays one pool dispatch, not five.  The explicit
1423    // fallback keeps mixed-dtype checkpoints on their established kernels.
1424    let q4_group = [&w.q_proj, &w.k_proj, &w.v_proj, &w.f_a, &w.b_proj];
1425    let q4_uniform = q4_group.iter().all(|t| {
1426        matches!(
1427            t,
1428            QTensor::Mapped {
1429                dtype: TensorDtype::Q4TiledP,
1430                ..
1431            }
1432        )
1433    });
1434    if q4_uniform {
1435        QTensor::matvec_many(
1436            q4_group,
1437            x,
1438            [
1439                xq.as_mut_slice(),
1440                xk.as_mut_slice(),
1441                xv.as_mut_slice(),
1442                fl.as_mut_slice(),
1443                b.as_mut_slice(),
1444            ],
1445            pool,
1446        );
1447    } else {
1448        QTensor::matvec_many(
1449            [&w.q_proj, &w.k_proj, &w.v_proj, &w.f_a],
1450            x,
1451            [
1452                xq.as_mut_slice(),
1453                xk.as_mut_slice(),
1454                xv.as_mut_slice(),
1455                fl.as_mut_slice(),
1456            ],
1457            pool,
1458        );
1459        w.b_proj.matvec(x, &mut b, pool);
1460    }
1461    let mut f = vec![0.0f32; kd];
1462    w.f_b.matvec(&fl, &mut f, pool);
1463    let gate_out = kda_gate_out(w, x, vd, pool);
1464
1465    let mut of = vec![0.0f32; vd];
1466    kda_step(
1467        &xq, &xk, &xv, &f, &b, &gate_out, w, cfg, state, &mut of, pool,
1468    );
1469
1470    let mut out = vec![0.0f32; cfg.hidden_size];
1471    w.o_proj.matvec(&of, &mut out, pool);
1472    out
1473}
1474
1475/// Batched KDA forward (prefill-GEMM): projections as matmat over the
1476/// chunk, the recurrence sequential per position — elementwise identical
1477/// to the single-position path.
1478pub fn kda_forward_batch(
1479    xs: &[f32],
1480    bsz: usize,
1481    w: &KdaWeights,
1482    cfg: &KdaCfg,
1483    state: &mut Vec<f32>,
1484    pool: Option<&Pool>,
1485) -> Vec<f32> {
1486    if state.len() != cfg.state_len() {
1487        *state = vec![0f32; cfg.state_len()];
1488    }
1489    let (nh, dk, dv, hs) = (
1490        cfg.num_heads,
1491        cfg.head_k_dim,
1492        cfg.head_v_dim,
1493        cfg.hidden_size,
1494    );
1495    let (kd, vd) = (nh * dk, nh * dv);
1496
1497    let mut xq = vec![0.0f32; bsz * kd];
1498    w.q_proj.matmat(xs, bsz, &mut xq, pool);
1499    let mut xk = vec![0.0f32; bsz * kd];
1500    w.k_proj.matmat(xs, bsz, &mut xk, pool);
1501    let mut xv = vec![0.0f32; bsz * vd];
1502    w.v_proj.matmat(xs, bsz, &mut xv, pool);
1503    let rank = w.f_a.rows();
1504    let mut fl = vec![0.0f32; bsz * rank];
1505    w.f_a.matmat(xs, bsz, &mut fl, pool);
1506    let mut f = vec![0.0f32; bsz * kd];
1507    w.f_b.matmat(&fl, bsz, &mut f, pool);
1508    let mut b = vec![0.0f32; bsz * nh];
1509    w.b_proj.matmat(xs, bsz, &mut b, pool);
1510    let mut gate_out = vec![0.0f32; bsz * vd];
1511    match &w.gate {
1512        KdaOutGate::Full(gp) => gp.matmat(xs, bsz, &mut gate_out, pool),
1513        KdaOutGate::LowRank(ga, gb) => {
1514            let mut low = vec![0.0f32; bsz * ga.rows()];
1515            ga.matmat(xs, bsz, &mut low, pool);
1516            gb.matmat(&low, bsz, &mut gate_out, pool);
1517        }
1518    }
1519
1520    let mut of = vec![0.0f32; bsz * vd];
1521    for bi in 0..bsz {
1522        let mut oh = vec![0.0f32; vd];
1523        kda_step(
1524            &xq[bi * kd..(bi + 1) * kd],
1525            &xk[bi * kd..(bi + 1) * kd],
1526            &xv[bi * vd..(bi + 1) * vd],
1527            &f[bi * kd..(bi + 1) * kd],
1528            &b[bi * nh..(bi + 1) * nh],
1529            &gate_out[bi * vd..(bi + 1) * vd],
1530            w,
1531            cfg,
1532            state,
1533            &mut oh,
1534            pool,
1535        );
1536        of[bi * vd..(bi + 1) * vd].copy_from_slice(&oh);
1537    }
1538
1539    let mut out = vec![0.0f32; bsz * hs];
1540    w.o_proj.matmat(&of, bsz, &mut out, pool);
1541    out
1542}
1543
1544#[cfg(test)]
1545mod tests {
1546    #[test]
1547    fn kda_forward_matches_naive_reference() {
1548        // Small deterministic KDA layer; the oracle is a literal port of
1549        // FLA naive_recurrent_kda + naive_kda_gate + the modeling glue
1550        // (conv→silu, low-rank decay, sigmoid-gated output norm), coded
1551        // straight from the reference — a different shape from the fused
1552        // two-pass production kernel.
1553        let (nh, dk, dv, kk, hs, rank) = (2usize, 4usize, 4usize, 3usize, 6usize, 3usize);
1554        let synth = |rows: usize, cols: usize, salt: usize| -> QTensor {
1555            QTensor::from_f32(
1556                (0..rows * cols)
1557                    .map(|i| (((i * 31 + salt * 17) % 101) as f32 / 101.0 - 0.5) * 0.6)
1558                    .collect(),
1559                rows,
1560                cols,
1561            )
1562        };
1563        let vecf = |n: usize, salt: usize| -> Vec<f32> {
1564            (0..n)
1565                .map(|i| (((i * 13 + salt * 7) % 89) as f32 / 89.0 - 0.5) * 0.8)
1566                .collect()
1567        };
1568        for (label, a_log, lb) in [
1569            ("per-head standard", vecf(nh, 40), None),
1570            ("per-dim lower-bound", vecf(dk, 41), Some(-5.0f32)),
1571        ] {
1572            let w = KdaWeights {
1573                q_proj: synth(nh * dk, hs, 1),
1574                k_proj: synth(nh * dk, hs, 2),
1575                v_proj: synth(nh * dv, hs, 3),
1576                conv_q: vecf(nh * dk * kk, 4),
1577                conv_k: vecf(nh * dk * kk, 5),
1578                conv_v: vecf(nh * dv * kk, 6),
1579                f_a: synth(rank, hs, 7),
1580                f_b: synth(nh * dk, rank, 8),
1581                dt_bias: vecf(nh * dk, 9),
1582                a_log: a_log.clone(),
1583                b_proj: synth(nh, hs, 10),
1584                gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
1585                o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
1586                o_proj: synth(hs, nh * dv, 13),
1587                gate_lower_bound: lb,
1588            };
1589            let cfg = KdaCfg {
1590                num_heads: nh,
1591                head_k_dim: dk,
1592                head_v_dim: dv,
1593                conv_kernel: kk,
1594                hidden_size: hs,
1595                rms_eps: 1e-6,
1596            };
1597            let xs: Vec<Vec<f32>> = (0..6)
1598                .map(|t| {
1599                    (0..hs)
1600                        .map(|i| ((t * hs + i) as f32 * 0.37).sin() * 0.5)
1601                        .collect()
1602                })
1603                .collect();
1604
1605            // Production path.
1606            let mut state = Vec::new();
1607            let got: Vec<Vec<f32>> = xs
1608                .iter()
1609                .map(|x| kda_forward(x, &w, &cfg, &mut state, None))
1610                .collect();
1611
1612            // Oracle.
1613            let mv = |t: &QTensor, x: &[f32]| -> Vec<f32> {
1614                let mut o = vec![0.0f32; t.rows()];
1615                t.matvec(x, &mut o, None);
1616                o
1617            };
1618            let mut hist: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = Vec::new(); // raw xq/xk/xv
1619            let mut s_state = vec![0f64; nh * dk * dv];
1620            let mut want: Vec<Vec<f32>> = Vec::new();
1621            for x in &xs {
1622                let (xq, xk, xv) = (mv(&w.q_proj, x), mv(&w.k_proj, x), mv(&w.v_proj, x));
1623                hist.push((xq, xk, xv));
1624                // conv over the raw history, taps oldest→newest.
1625                let conv = |sel: fn(&(Vec<f32>, Vec<f32>, Vec<f32>)) -> &Vec<f32>,
1626                            taps: &[f32],
1627                            n: usize|
1628                 -> Vec<f32> {
1629                    (0..n)
1630                        .map(|c| {
1631                            let t = &taps[c * kk..(c + 1) * kk];
1632                            let mut acc = 0f64;
1633                            for j in 0..kk {
1634                                let idx = hist.len() as i64 - (kk as i64 - j as i64);
1635                                if idx >= 0 {
1636                                    acc += sel(&hist[idx as usize])[c] as f64 * t[j] as f64;
1637                                }
1638                            }
1639                            silu(acc)
1640                        })
1641                        .map(|v| v as f32)
1642                        .collect()
1643                };
1644                let cq = conv(|h| &h.0, &w.conv_q, nh * dk);
1645                let ck = conv(|h| &h.1, &w.conv_k, nh * dk);
1646                let cv = conv(|h| &h.2, &w.conv_v, nh * dv);
1647                let f = mv(&w.f_b, &mv(&w.f_a, x));
1648                let bb = mv(&w.b_proj, x);
1649                let gate_out = match &w.gate {
1650                    KdaOutGate::LowRank(ga, gb) => mv(gb, &mv(ga, x)),
1651                    KdaOutGate::Full(g) => mv(g, x),
1652                };
1653                let mut of = vec![0f32; nh * dv];
1654                for h in 0..nh {
1655                    // l2norm + scale.
1656                    let q: Vec<f64> = {
1657                        let sl = &cq[h * dk..(h + 1) * dk];
1658                        let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
1659                        let inv = 1.0 / ((n + 1e-6).sqrt() * (dk as f64).sqrt());
1660                        sl.iter().map(|&v| v as f64 * inv).collect()
1661                    };
1662                    let k: Vec<f64> = {
1663                        let sl = &ck[h * dk..(h + 1) * dk];
1664                        let n: f64 = sl.iter().map(|&v| (v as f64) * (v as f64)).sum();
1665                        let inv = 1.0 / (n + 1e-6).sqrt();
1666                        sl.iter().map(|&v| v as f64 * inv).collect()
1667                    };
1668                    let v: Vec<f64> = cv[h * dv..(h + 1) * dv].iter().map(|&v| v as f64).collect();
1669                    // gate: g = −exp(A)·softplus(f+bias) | lb·σ(exp(A)·(f+bias))
1670                    let g: Vec<f64> = (0..dk)
1671                        .map(|d| {
1672                            let a = if w.a_log.len() == nh {
1673                                w.a_log[h] as f64
1674                            } else {
1675                                w.a_log[d] as f64
1676                            };
1677                            let raw = f[h * dk + d] as f64 + w.dt_bias[h * dk + d] as f64;
1678                            match w.gate_lower_bound {
1679                                Some(lb) => lb as f64 * sigmoid(a.exp() * raw),
1680                                None => -a.exp() * softplus(raw),
1681                            }
1682                        })
1683                        .collect();
1684                    let beta = sigmoid(bb[h] as f64);
1685                    let s = &mut s_state[h * dk * dv..(h + 1) * dk * dv];
1686                    // S = Diag(exp(g))·S
1687                    for di in 0..dk {
1688                        for dj in 0..dv {
1689                            s[di * dv + dj] *= g[di].exp();
1690                        }
1691                    }
1692                    // kv = kᵀS; S += β·k⊗(v−kv); o = qᵀS
1693                    let mut kv = vec![0f64; dv];
1694                    for di in 0..dk {
1695                        for dj in 0..dv {
1696                            kv[dj] += k[di] * s[di * dv + dj];
1697                        }
1698                    }
1699                    for di in 0..dk {
1700                        for dj in 0..dv {
1701                            s[di * dv + dj] += beta * k[di] * (v[dj] - kv[dj]);
1702                        }
1703                    }
1704                    let mut o = vec![0f64; dv];
1705                    for di in 0..dk {
1706                        for dj in 0..dv {
1707                            o[dj] += q[di] * s[di * dv + dj];
1708                        }
1709                    }
1710                    // sigmoid-gated RMSNorm
1711                    let ss: f64 = o.iter().map(|&v| v * v).sum();
1712                    let inv = 1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt();
1713                    for dj in 0..dv {
1714                        of[h * dv + dj] = (o[dj]
1715                            * inv
1716                            * w.o_norm[dj] as f64
1717                            * sigmoid(gate_out[h * dv + dj] as f64))
1718                            as f32;
1719                    }
1720                }
1721                want.push(mv(&w.o_proj, &of));
1722            }
1723
1724            for (t, (g, e)) in got.iter().zip(&want).enumerate() {
1725                for (i, (a, b)) in g.iter().zip(e.iter()).enumerate() {
1726                    assert!((a - b).abs() < 2e-4, "{label}: t={t} i={i}: {a} vs {b}");
1727                }
1728            }
1729        }
1730
1731        // Batched prefill must equal the sequential singles bit-close.
1732        let w = KdaWeights {
1733            q_proj: synth(nh * dk, hs, 1),
1734            k_proj: synth(nh * dk, hs, 2),
1735            v_proj: synth(nh * dv, hs, 3),
1736            conv_q: vecf(nh * dk * kk, 4),
1737            conv_k: vecf(nh * dk * kk, 5),
1738            conv_v: vecf(nh * dv * kk, 6),
1739            f_a: synth(rank, hs, 7),
1740            f_b: synth(nh * dk, rank, 8),
1741            dt_bias: vecf(nh * dk, 9),
1742            a_log: vecf(nh, 40),
1743            b_proj: synth(nh, hs, 10),
1744            gate: KdaOutGate::LowRank(synth(rank, hs, 11), synth(nh * dv, rank, 12)),
1745            o_norm: (0..dv).map(|i| 1.0 + 0.1 * i as f32).collect(),
1746            o_proj: synth(hs, nh * dv, 13),
1747            gate_lower_bound: None,
1748        };
1749        let cfg = KdaCfg {
1750            num_heads: nh,
1751            head_k_dim: dk,
1752            head_v_dim: dv,
1753            conv_kernel: kk,
1754            hidden_size: hs,
1755            rms_eps: 1e-6,
1756        };
1757        let xs: Vec<f32> = (0..5 * hs).map(|i| (i as f32 * 0.29).cos() * 0.4).collect();
1758        let mut st1 = Vec::new();
1759        let seq: Vec<f32> = (0..5)
1760            .flat_map(|t| kda_forward(&xs[t * hs..(t + 1) * hs], &w, &cfg, &mut st1, None))
1761            .collect();
1762        let mut st2 = Vec::new();
1763        let bat = kda_forward_batch(&xs, 5, &w, &cfg, &mut st2, None);
1764        for (i, (a, b)) in seq.iter().zip(&bat).enumerate() {
1765            assert!((a - b).abs() < 1e-5, "batch i={i}: {a} vs {b}");
1766        }
1767        assert_eq!(st1, st2, "state must match after the chunk");
1768    }
1769
1770    use super::*;
1771
1772    fn tiny() -> (VmfPhaseWeights, VmfPhaseCfg) {
1773        let cfg = VmfPhaseCfg {
1774            num_heads: 2,
1775            nphase: 3,
1776            value_head_dim: 4,
1777            hidden_size: 8,
1778            phase_mass: 0.0,
1779        };
1780        let synth = |rows: usize, cols: usize, salt: usize| {
1781            QTensor::from_f32(
1782                (0..rows * cols)
1783                    .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
1784                    .collect(),
1785                rows,
1786                cols,
1787            )
1788        };
1789        let w = VmfPhaseWeights {
1790            thq: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 1),
1791            thk: synth(cfg.num_heads * cfg.nphase, cfg.hidden_size, 2),
1792            v_proj: synth(cfg.num_heads * cfg.value_head_dim, cfg.hidden_size, 3),
1793            out_proj: synth(cfg.hidden_size, cfg.num_heads * cfg.value_head_dim, 4),
1794            decay: (0..cfg.num_heads * 2 * cfg.nphase)
1795                .map(|i| 0.9 + 0.005 * (i % 10) as f64)
1796                .collect(),
1797            conv: None,
1798            k_gate: None,
1799            phase_delta: false,
1800        };
1801        (w, cfg)
1802    }
1803
1804    #[test]
1805    fn state_persists_and_changes_output() {
1806        let (w, cfg) = tiny();
1807        let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
1808        let mut state = Vec::new();
1809        let o1 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1810        let o2 = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1811        // Same input, evolved state → different output.
1812        assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
1813        assert_eq!(state.len(), cfg.state_len());
1814    }
1815
1816    #[test]
1817    fn phase_delta_matches_independent_token_oracle_and_ignores_phase_mass() {
1818        // One head, two phase angles, one value channel.  The projections are
1819        // intentionally simple so this test computes the token oracle without
1820        // calling any implementation helper: q=[x0,x1], k=[x1,x0], v=x0,
1821        // and the first output row is qᵀS.
1822        let cfg = VmfPhaseCfg {
1823            num_heads: 1,
1824            nphase: 2,
1825            value_head_dim: 1,
1826            hidden_size: 2,
1827            // A nonzero legacy knob must not alter a v1 operator.
1828            phase_mass: 17.0,
1829        };
1830        let w = VmfPhaseWeights {
1831            thq: QTensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], 2, 2),
1832            thk: QTensor::from_f32(vec![0.0, 1.0, 1.0, 0.0], 2, 2),
1833            v_proj: QTensor::from_f32(vec![1.0, 0.0], 1, 2),
1834            out_proj: QTensor::from_f32(vec![1.0, 0.0], 2, 1),
1835            decay: vec![0.9, 0.8, 0.7, 0.6],
1836            conv: None,
1837            k_gate: None, // beta=1
1838            phase_delta: true,
1839        };
1840        let xs = [[0.3f32, 0.7f32], [-0.2, 0.4], [0.8, -0.6]];
1841        let c = 1.0f64 / 2.0f64.sqrt();
1842        let mut oracle_state = [0.0f64; 4];
1843        let mut state = Vec::new();
1844        for x in xs {
1845            let q = [
1846                c * (x[0] as f64).cos(),
1847                c * (x[1] as f64).cos(),
1848                c * (x[0] as f64).sin(),
1849                c * (x[1] as f64).sin(),
1850            ];
1851            let k = [
1852                c * (x[1] as f64).cos(),
1853                c * (x[0] as f64).cos(),
1854                c * (x[1] as f64).sin(),
1855                c * (x[0] as f64).sin(),
1856            ];
1857            let value = x[0] as f64;
1858            let mut read = 0.0;
1859            for f in 0..4 {
1860                oracle_state[f] *= w.decay[f];
1861                read += k[f] * oracle_state[f];
1862            }
1863            for f in 0..4 {
1864                oracle_state[f] += k[f] * (value - read);
1865            }
1866            let want: f32 = q
1867                .iter()
1868                .zip(oracle_state)
1869                .map(|(qf, sf)| (qf * sf) as f32)
1870                .sum();
1871            let got = vmf_phase_forward(&x, &w, &cfg, &mut state, None);
1872            assert!(
1873                (got[0] - want).abs() < 2e-6,
1874                "oracle {want} vs runtime {}",
1875                got[0]
1876            );
1877            assert!(got[1].abs() < 2e-7);
1878        }
1879        assert_eq!(state.len(), cfg.state_len());
1880        for (got, want) in state.iter().zip(oracle_state) {
1881            assert!((*got as f64 - want).abs() < 2e-6);
1882        }
1883    }
1884
1885    #[test]
1886    fn phase_delta_fused_readout_is_bit_exact_over_long_continuation() {
1887        // Independent three-pass reference, with nonzero state/output and a
1888        // sentinel convolution tail. Exercise odd widths and the real shape.
1889        for (nh, nph, dv) in [(1, 1, 1), (3, 5, 7), (8, 32, 128)] {
1890            let cfg = VmfPhaseCfg {
1891                num_heads: nh,
1892                nphase: nph,
1893                value_head_dim: dv,
1894                hidden_size: nh * dv,
1895                phase_mass: 19.0,
1896            };
1897            let p2 = 2 * nph;
1898            let scale = 1.0f64 / (nph as f64).sqrt();
1899            let mut state: Vec<f32> = (0..cfg.state_len() + 11)
1900                .map(|i| (i as f32 * 0.37).sin())
1901                .collect();
1902            let mut reference = state.clone();
1903            let tail = state[cfg.state_len()..].to_vec();
1904            let decay: Vec<f64> = (0..nh * p2).map(|i| [0.0, 0.8, 0.99, 1.0][i % 4]).collect();
1905            for step in 0..128 {
1906                let q: Vec<f32> = (0..nh * nph)
1907                    .map(|i| ((i + step * 3) as f32 * 0.13).sin() * 4.0)
1908                    .collect();
1909                let k: Vec<f32> = (0..nh * nph)
1910                    .map(|i| ((i + step * 7) as f32 * 0.29).cos() * 4.0)
1911                    .collect();
1912                let v: Vec<f32> = (0..nh * dv)
1913                    .map(|i| ((i + step) as f32 * 0.43).sin())
1914                    .collect();
1915                let gates: Vec<f32> = (0..nh).map(|h| [0.0, 0.3, 1.0][(step + h) % 3]).collect();
1916                let kap = (step % 4 != 0).then_some(gates.as_slice());
1917                let mut out = vec![0.125f32; nh * dv];
1918                let mut want = out.clone();
1919                for h in 0..nh {
1920                    let feature = |theta: &[f32], f: usize| {
1921                        let angle = theta[h * nph + f % nph] as f64;
1922                        scale * if f < nph { angle.cos() } else { angle.sin() }
1923                    };
1924                    let beta = kap.map_or(1.0f64, |g| g[h] as f64);
1925                    let mut read = vec![0.0f64; dv];
1926                    for f in 0..p2 {
1927                        for d in 0..dv {
1928                            let at = (h * p2 + f) * dv + d;
1929                            read[d] += feature(&k, f) * (decay[h * p2 + f] * reference[at] as f64);
1930                        }
1931                    }
1932                    for f in 0..p2 {
1933                        for d in 0..dv {
1934                            let at = (h * p2 + f) * dv + d;
1935                            reference[at] = (decay[h * p2 + f] * reference[at] as f64
1936                                + beta * feature(&k, f) * (v[h * dv + d] as f64 - read[d]))
1937                                as f32;
1938                        }
1939                    }
1940                    for f in 0..p2 {
1941                        for d in 0..dv {
1942                            want[h * dv + d] +=
1943                                (feature(&q, f) * reference[(h * p2 + f) * dv + d] as f64) as f32;
1944                        }
1945                    }
1946                }
1947                phase_delta_step_f32(&q, &k, &v, &decay, kap, &cfg, &mut state, &mut out);
1948                assert_eq!(out, want, "shape {nh}/{nph}/{dv}, step {step}");
1949                assert_eq!(state, reference, "state at step {step}");
1950                assert_eq!(&state[cfg.state_len()..], tail);
1951            }
1952        }
1953    }
1954
1955    #[test]
1956    fn phase_delta_pair_reset_and_legacy_paths_are_distinct() {
1957        let (mut delta, mut cfg) = tiny();
1958        delta.phase_delta = true;
1959        // The phase normalization is observable for nphase=3, while the
1960        // legacy core remains unchanged and still uses raw cos/sin features.
1961        let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
1962        let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
1963        let mut seq_state = Vec::new();
1964        let d1 = vmf_phase_forward(&x1, &delta, &cfg, &mut seq_state, None);
1965        let d2 = vmf_phase_forward(&x2, &delta, &cfg, &mut seq_state, None);
1966
1967        let mut pair_state = Vec::new();
1968        let mut scratch = Vec::new();
1969        let (p1, p2) = vmf_phase_pair(&x1, &x2, &delta, &cfg, &mut pair_state, &mut scratch, None);
1970        assert_eq!(d1, p1);
1971        assert_eq!(d2, p2);
1972        std::mem::swap(&mut pair_state, &mut scratch);
1973        assert_eq!(seq_state, pair_state);
1974
1975        // Empty state is the runtime reset contract; it reproduces the first
1976        // token exactly and does not inherit the previous sequence.
1977        let mut reset = seq_state;
1978        reset.clear();
1979        let after_reset = vmf_phase_forward(&x1, &delta, &cfg, &mut reset, None);
1980        assert_eq!(after_reset, d1);
1981
1982        cfg.phase_mass = 0.0;
1983        let mut legacy = delta;
1984        legacy.phase_delta = false;
1985        let mut legacy_state = Vec::new();
1986        let old = vmf_phase_forward(&x1, &legacy, &cfg, &mut legacy_state, None);
1987        assert!(
1988            old.iter().zip(&d1).any(|(a, b)| (a - b).abs() > 1e-5),
1989            "legacy additive and normalized Phase-Delta paths must differ"
1990        );
1991    }
1992
1993    /// Phase-mass correction: mass=0 is bit-identical to the unscaled kernel; mass>0
1994    /// changes the output (phase narrowed → kernel widened). Guards the
1995    /// no-op default and that the knob is actually wired.
1996    #[test]
1997    fn phase_mass_zero_is_noop_and_positive_shifts() {
1998        let (w, cfg0) = tiny();
1999        let mut cfg_m = cfg0;
2000        cfg_m.phase_mass = 1.0;
2001        let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.4).sin()).collect();
2002
2003        let mut s0 = Vec::new();
2004        let base = vmf_phase_forward(&x, &w, &cfg0, &mut s0, None);
2005        // Re-run with mass=0 → must be bit-identical.
2006        let mut s0b = Vec::new();
2007        let base2 = vmf_phase_forward(&x, &w, &cfg0, &mut s0b, None);
2008        assert_eq!(base, base2, "mass=0 must be deterministic/no-op");
2009        // mass=1 → output differs (θ halved before cos/sin).
2010        let mut sm = Vec::new();
2011        let massed = vmf_phase_forward(&x, &w, &cfg_m, &mut sm, None);
2012        assert!(
2013            base.iter().zip(&massed).any(|(a, b)| (a - b).abs() > 1e-5),
2014            "mass>0 must change the output"
2015        );
2016        assert!(massed.iter().all(|v| v.is_finite()));
2017    }
2018
2019    /// κ write gate (hybrid_k): saturated-open gate (bias ≫ 0 → κ→1)
2020    /// matches the gateless kernel within fp tolerance; a closed gate
2021    /// (bias ≪ 0 → κ→0) writes nothing — the state stays zero and the
2022    /// output collapses to the empty-state readout.
2023    #[test]
2024    fn kappa_gate_open_matches_none_and_closed_writes_nothing() {
2025        let (mut w, cfg) = tiny();
2026        let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
2027
2028        let mut s_none = Vec::new();
2029        let base1 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
2030        let base2 = vmf_phase_forward(&x, &w, &cfg, &mut s_none, None);
2031
2032        // Open gate: W=0, bias=+20 → κ = σ(20) ≈ 1 − 2e−9.
2033        w.k_gate = Some((
2034            QTensor::from_f32(
2035                vec![0.0; cfg.num_heads * cfg.hidden_size],
2036                cfg.num_heads,
2037                cfg.hidden_size,
2038            ),
2039            vec![20.0; cfg.num_heads],
2040        ));
2041        let mut s_open = Vec::new();
2042        let o1 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
2043        let o2 = vmf_phase_forward(&x, &w, &cfg, &mut s_open, None);
2044        for (a, b) in base1.iter().zip(&o1).chain(base2.iter().zip(&o2)) {
2045            assert!(
2046                (a - b).abs() < 1e-5,
2047                "open κ must match gateless: {a} vs {b}"
2048            );
2049        }
2050
2051        // Closed gate: bias=−20 → κ ≈ 0 → nothing is written.
2052        w.k_gate = Some((
2053            QTensor::from_f32(
2054                vec![0.0; cfg.num_heads * cfg.hidden_size],
2055                cfg.num_heads,
2056                cfg.hidden_size,
2057            ),
2058            vec![-20.0; cfg.num_heads],
2059        ));
2060        let mut s_closed = Vec::new();
2061        let oc = vmf_phase_forward(&x, &w, &cfg, &mut s_closed, None);
2062        assert!(
2063            s_closed.iter().all(|&v| v.abs() < 1e-7),
2064            "closed κ: state must stay empty"
2065        );
2066        assert!(
2067            oc.iter().all(|&v| v.abs() < 1e-6),
2068            "closed κ: empty-state readout"
2069        );
2070    }
2071
2072    #[test]
2073    fn pair_matches_two_singles_bitexact() {
2074        let (w, cfg) = tiny();
2075        let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
2076        let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
2077
2078        // Reference: two sequential singles.
2079        let mut s_ref = Vec::new();
2080        let r1 = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
2081        let r2 = vmf_phase_forward(&x2, &w, &cfg, &mut s_ref, None);
2082
2083        // Pair: lane1 commits, lane2 tentative in scratch.
2084        let mut s = Vec::new();
2085        let mut scratch = Vec::new();
2086        let (p1, p2) = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2087        assert_eq!(r1, p1, "lane 1 must be bit-identical");
2088        assert_eq!(r2, p2, "lane 2 must be bit-identical");
2089        // Accepting the draft = swapping scratch in → equals s_ref.
2090        std::mem::swap(&mut s, &mut scratch);
2091        assert_eq!(s, s_ref, "accepted state must equal sequential state");
2092    }
2093
2094    #[test]
2095    fn rejected_draft_leaves_state_at_lane1() {
2096        let (w, cfg) = tiny();
2097        let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
2098        let x2 = vec![0.5f32; 8];
2099
2100        let mut s_ref = Vec::new();
2101        let _ = vmf_phase_forward(&x1, &w, &cfg, &mut s_ref, None);
2102
2103        let mut s = Vec::new();
2104        let mut scratch = Vec::new();
2105        let _ = vmf_phase_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2106        // Reject: state must be exactly the post-lane1 state.
2107        assert_eq!(s, s_ref);
2108    }
2109
2110    // ───────────── GatedDeltaNet ─────────────
2111
2112    fn tiny_gdn() -> (GdnWeights, GdnCfg) {
2113        let cfg = GdnCfg {
2114            num_v_heads: 4,
2115            num_k_heads: 2,
2116            key_head_dim: 3,
2117            value_head_dim: 5,
2118            conv_kernel: 4,
2119            hidden_size: 8,
2120            rms_eps: 1e-6,
2121            output_gate_sigmoid: false,
2122        };
2123        let c_dim = cfg.conv_dim();
2124        let vd = cfg.num_v_heads * cfg.value_head_dim;
2125        let synth = |rows: usize, cols: usize, salt: usize| {
2126            QTensor::from_f32(
2127                (0..rows * cols)
2128                    .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
2129                    .collect(),
2130                rows,
2131                cols,
2132            )
2133        };
2134        let vecf = |n: usize, salt: usize| -> Vec<f32> {
2135            (0..n)
2136                .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.6)
2137                .collect()
2138        };
2139        let w = GdnWeights {
2140            in_proj_qkv: synth(c_dim, cfg.hidden_size, 1),
2141            in_proj_z: synth(vd, cfg.hidden_size, 2),
2142            in_proj_a: synth(cfg.num_v_heads, cfg.hidden_size, 3),
2143            in_proj_b: synth(cfg.num_v_heads, cfg.hidden_size, 4),
2144            conv1d: vecf(c_dim * cfg.conv_kernel, 5),
2145            a_log: (0..cfg.num_v_heads).map(|i| 0.2 + 0.3 * i as f32).collect(),
2146            dt_bias: vecf(cfg.num_v_heads, 6),
2147            norm: vec![1.0; cfg.value_head_dim],
2148            out_proj: synth(cfg.hidden_size, vd, 7),
2149        };
2150        (w, cfg)
2151    }
2152
2153    #[test]
2154    fn gdn_state_persists_and_changes_output() {
2155        let (w, cfg) = tiny_gdn();
2156        let x: Vec<f32> = (0..8).map(|i| (i as f32 * 0.3).sin()).collect();
2157        let mut state = Vec::new();
2158        let o1 = gdn_forward(&x, &w, &cfg, &mut state, None);
2159        let o2 = gdn_forward(&x, &w, &cfg, &mut state, None);
2160        assert!(o1.iter().zip(&o2).any(|(a, b)| (a - b).abs() > 1e-6));
2161        assert_eq!(state.len(), cfg.state_len());
2162    }
2163
2164    #[test]
2165    fn gdn_pair_matches_two_singles_bitexact() {
2166        let (w, cfg) = tiny_gdn();
2167        let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.2).cos()).collect();
2168        let x2: Vec<f32> = (0..8).map(|i| (i as f32 * 0.5).sin()).collect();
2169
2170        let mut s_ref = Vec::new();
2171        let r1 = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
2172        let r2 = gdn_forward(&x2, &w, &cfg, &mut s_ref, None);
2173
2174        let mut s = Vec::new();
2175        let mut scratch = Vec::new();
2176        let (p1, p2) = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2177        assert_eq!(r1, p1, "lane 1 must be bit-identical");
2178        assert_eq!(r2, p2, "lane 2 must be bit-identical");
2179        std::mem::swap(&mut s, &mut scratch);
2180        assert_eq!(s, s_ref, "accepted state must equal sequential state");
2181    }
2182
2183    #[test]
2184    fn gdn_rejected_draft_leaves_state_at_lane1() {
2185        let (w, cfg) = tiny_gdn();
2186        let x1: Vec<f32> = (0..8).map(|i| (i as f32 * 0.7).sin()).collect();
2187        let x2 = vec![0.5f32; 8];
2188
2189        let mut s_ref = Vec::new();
2190        let _ = gdn_forward(&x1, &w, &cfg, &mut s_ref, None);
2191
2192        let mut s = Vec::new();
2193        let mut scratch = Vec::new();
2194        let _ = gdn_pair(&x1, &x2, &w, &cfg, &mut s, &mut scratch, None);
2195        assert_eq!(s, s_ref);
2196    }
2197
2198    /// The conv ring must give the same result as an explicit causal
2199    /// conv over the whole sequence (oracle semantics: zero left-pad,
2200    /// tap kk−1 on the current position).
2201    #[test]
2202    fn gdn_conv_ring_matches_explicit_causal_conv() {
2203        let (w, cfg) = tiny_gdn();
2204        let seq: Vec<Vec<f32>> = (0..6)
2205            .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.17).sin()).collect())
2206            .collect();
2207
2208        // Reference: recompute position t from scratch each time with a
2209        // fresh state built by replaying the prefix.
2210        let mut s_inc = Vec::new();
2211        for (t, x) in seq.iter().enumerate() {
2212            let inc = gdn_forward(x, &w, &cfg, &mut s_inc, None);
2213            let mut s_replay = Vec::new();
2214            let mut replay = Vec::new();
2215            for xr in &seq[..=t] {
2216                replay = gdn_forward(xr, &w, &cfg, &mut s_replay, None);
2217            }
2218            assert_eq!(inc, replay, "position {t}: ring must equal replay");
2219        }
2220    }
2221
2222    fn tiny_short_conv() -> (ShortConvWeights, ShortConvCfg) {
2223        let cfg = ShortConvCfg {
2224            hidden_size: 8,
2225            kernel: 3,
2226        };
2227        let synth = |rows: usize, cols: usize, salt: usize| {
2228            QTensor::from_f32(
2229                (0..rows * cols)
2230                    .map(|i| (((i * 11 + salt * 5) % 89) as f32 / 89.0 - 0.5) * 0.5)
2231                    .collect(),
2232                rows,
2233                cols,
2234            )
2235        };
2236        let w = ShortConvWeights {
2237            in_proj: synth(3 * cfg.hidden_size, cfg.hidden_size, 1),
2238            conv: (0..cfg.hidden_size * cfg.kernel)
2239                .map(|i| ((i * 7 % 13) as f32 / 13.0 - 0.5) * 0.8)
2240                .collect(),
2241            out_proj: synth(cfg.hidden_size, cfg.hidden_size, 2),
2242        };
2243        (w, cfg)
2244    }
2245
2246    /// The incremental conv ring must equal a from-scratch causal replay
2247    /// of the prefix at every position — the decode/prefill contract.
2248    #[test]
2249    fn short_conv_ring_matches_explicit_causal_conv() {
2250        let (w, cfg) = tiny_short_conv();
2251        let seq: Vec<Vec<f32>> = (0..6)
2252            .map(|t| (0..8).map(|i| ((t * 8 + i) as f32 * 0.19).cos()).collect())
2253            .collect();
2254        let mut s_inc = Vec::new();
2255        for (t, x) in seq.iter().enumerate() {
2256            let inc = short_conv_forward(x, &w, &cfg, &mut s_inc, None);
2257            let mut s_replay = Vec::new();
2258            let mut replay = Vec::new();
2259            for xr in &seq[..=t] {
2260                replay = short_conv_forward(xr, &w, &cfg, &mut s_replay, None);
2261            }
2262            assert_eq!(inc, replay, "position {t}: ring must equal replay");
2263            assert_eq!(s_inc.len(), cfg.state_len());
2264        }
2265    }
2266
2267    /// The batched prefill path (matmat + sequential conv over the chunk)
2268    /// must reproduce the position-by-position decode path exactly.
2269    #[test]
2270    fn short_conv_batch_matches_sequential() {
2271        let (w, cfg) = tiny_short_conv();
2272        let b = 5;
2273        let xs: Vec<f32> = (0..b * cfg.hidden_size)
2274            .map(|i| (i as f32 * 0.13).sin() * 0.6)
2275            .collect();
2276
2277        let mut s_seq = Vec::new();
2278        let mut seq_out = vec![0.0f32; b * cfg.hidden_size];
2279        for bi in 0..b {
2280            let o = short_conv_forward(
2281                &xs[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size],
2282                &w,
2283                &cfg,
2284                &mut s_seq,
2285                None,
2286            );
2287            seq_out[bi * cfg.hidden_size..(bi + 1) * cfg.hidden_size].copy_from_slice(&o);
2288        }
2289
2290        let mut s_batch = Vec::new();
2291        let batch_out = short_conv_forward_batch(&xs, b, &w, &cfg, &mut s_batch, None);
2292        assert_eq!(
2293            seq_out, batch_out,
2294            "batch conv must match sequential decode"
2295        );
2296        assert_eq!(s_seq, s_batch, "ring state must match after the chunk");
2297    }
2298}