Skip to main content

cortiq_engine/
fcd_ops.rs

1//! FCD polish — hand-rolled forward/backward operators.
2//!
3//! The training graph is FIXED (docs/RUST_FCD.md), so there is no
4//! autograd and no tape: every operator here is a (forward, backward)
5//! pair written out by hand, llm.c style. Ops are generic over the
6//! minimal `Fp` float trait so the SAME code that trains in f32 is
7//! gradchecked in f64 against central finite differences
8//! (tests/fcd_gradcheck.rs, rel err < 1e-3).
9//!
10//! Conventions:
11//! - all matrices are row-major; weight matrices use the runtime layout
12//!   `[out_dim, in_dim]` (a matvec is out[o] = dot(w_row_o, x));
13//! - every backward ACCUMULATES into its output gradient buffers
14//!   (`+=`) — callers zero them once per step, residual branches then
15//!   just add up naturally;
16//! - the attention-weight ops (`attn_head_*`, `nystrom_head_*`) are
17//!   meant to run in f64: the certified CPU probe computed the T×T
18//!   weight matrices in f64, which also makes raw `exp(±40)` safe with
19//!   no flash-shift bookkeeping.
20
21use crate::pool::Pool;
22
23// ─────────────────────────── float trait ───────────────────────────
24
25/// Minimal float abstraction: just enough for the fixed graph. Not a
26/// general numeric tower — two impls (f32/f64), no external crates.
27pub trait Fp:
28    Copy
29    + PartialOrd
30    + core::ops::Add<Output = Self>
31    + core::ops::Sub<Output = Self>
32    + core::ops::Mul<Output = Self>
33    + core::ops::Div<Output = Self>
34    + core::ops::Neg<Output = Self>
35    + core::ops::AddAssign
36    + core::ops::MulAssign
37    + Send
38    + Sync
39    + 'static
40{
41    const ZERO: Self;
42    const ONE: Self;
43    fn exp(self) -> Self;
44    fn sqrt(self) -> Self;
45    fn maxf(self, o: Self) -> Self;
46    fn fromf(x: f64) -> Self;
47    fn f64(self) -> f64;
48}
49
50impl Fp for f32 {
51    const ZERO: Self = 0.0;
52    const ONE: Self = 1.0;
53    #[inline]
54    fn exp(self) -> Self {
55        f32::exp(self)
56    }
57    #[inline]
58    fn sqrt(self) -> Self {
59        f32::sqrt(self)
60    }
61    #[inline]
62    fn maxf(self, o: Self) -> Self {
63        f32::max(self, o)
64    }
65    #[inline]
66    fn fromf(x: f64) -> Self {
67        x as f32
68    }
69    #[inline]
70    fn f64(self) -> f64 {
71        self as f64
72    }
73}
74
75impl Fp for f64 {
76    const ZERO: Self = 0.0;
77    const ONE: Self = 1.0;
78    #[inline]
79    fn exp(self) -> Self {
80        f64::exp(self)
81    }
82    #[inline]
83    fn sqrt(self) -> Self {
84        f64::sqrt(self)
85    }
86    #[inline]
87    fn maxf(self, o: Self) -> Self {
88        f64::max(self, o)
89    }
90    #[inline]
91    fn fromf(x: f64) -> Self {
92        x
93    }
94    #[inline]
95    fn f64(self) -> f64 {
96        self
97    }
98}
99
100#[inline]
101fn dot<F: Fp>(a: &[F], b: &[F]) -> F {
102    let mut s = F::ZERO;
103    for (x, y) in a.iter().zip(b) {
104        s += *x * *y;
105    }
106    s
107}
108
109// ─────────────────────── generic matmul (nt) ───────────────────────
110
111/// y[n,m] = x[n,k] · w[m,k]ᵀ (serial reference; the f32 hot path is
112/// `gemm_nt` below — same math, blocked + pooled).
113pub fn matmul_nt<F: Fp>(x: &[F], w: &[F], y: &mut [F], n: usize, k: usize, m: usize) {
114    for i in 0..n {
115        let xr = &x[i * k..(i + 1) * k];
116        for o in 0..m {
117            y[i * m + o] = dot(xr, &w[o * k..(o + 1) * k]);
118        }
119    }
120}
121
122/// dX += dY[n,m] · W[m,k] — the input-gradient half of matmul_nt.
123pub fn matmul_nt_dx<F: Fp>(dy: &[F], w: &[F], dx: &mut [F], n: usize, k: usize, m: usize) {
124    for i in 0..n {
125        let dxr = &mut dx[i * k..(i + 1) * k];
126        for o in 0..m {
127            let g = dy[i * m + o];
128            for (d, wv) in dxr.iter_mut().zip(&w[o * k..(o + 1) * k]) {
129                *d += g * *wv;
130            }
131        }
132    }
133}
134
135/// dW += dYᵀ[m,n] · X[n,k] — the weight-gradient half of matmul_nt.
136pub fn matmul_nt_dw<F: Fp>(dy: &[F], x: &[F], dw: &mut [F], n: usize, k: usize, m: usize) {
137    for i in 0..n {
138        let xr = &x[i * k..(i + 1) * k];
139        for o in 0..m {
140            let g = dy[i * m + o];
141            for (d, xv) in dw[o * k..(o + 1) * k].iter_mut().zip(xr) {
142                *d += g * *xv;
143            }
144        }
145    }
146}
147
148// ───────────────────── pooled f32 GEMM hot path ─────────────────────
149
150/// Row block: the block's X stays cache-resident while W streams once
151/// per block (a per-row W stream would move gigabytes per matmul).
152const GEMM_BLOCK: usize = 128;
153
154/// Pointer wrapper for disjoint parallel writes (same pattern as the
155/// pipeline's scatter — workers touch disjoint index ranges).
156struct SendMut<T>(*mut T);
157unsafe impl<T> Send for SendMut<T> {}
158unsafe impl<T> Sync for SendMut<T> {}
159impl<T> SendMut<T> {
160    #[inline]
161    // Deliberate unsynchronized scatter: pool workers write disjoint
162    // ranges in parallel, so returning `&mut` from `&self` is
163    // intentional here (same pattern as the pipeline's SendMut).
164    #[allow(clippy::mut_from_ref)]
165    unsafe fn slice(&self, off: usize, len: usize) -> &mut [T] {
166        unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
167    }
168}
169
170/// y[n,m] = x[n,k] · w[m,k]ᵀ — parallel over row blocks; bit-identical
171/// to `matmul_nt` (disjoint rows, same dot kernel regrouped by NEON).
172
173/// Accelerate CBLAS fast path for the training GEMMs (macOS): the
174/// naive pooled kernels below stay as the portable/reference path.
175#[cfg(target_os = "macos")]
176mod accel {
177    #[link(name = "Accelerate", kind = "framework")]
178    unsafe extern "C" {
179        pub fn cblas_sgemm(
180            order: i32,
181            ta: i32,
182            tb: i32,
183            m: i32,
184            n: i32,
185            k: i32,
186            alpha: f32,
187            a: *const f32,
188            lda: i32,
189            b: *const f32,
190            ldb: i32,
191            beta: f32,
192            c: *mut f32,
193            ldc: i32,
194        );
195    }
196
197    pub fn on() -> bool {
198        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
199        *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
200    }
201}
202
203pub fn gemm_nt(
204    x: &[f32],
205    w: &[f32],
206    y: &mut [f32],
207    n: usize,
208    k: usize,
209    m: usize,
210    pool: Option<&Pool>,
211) {
212    debug_assert_eq!(x.len(), n * k);
213    debug_assert_eq!(w.len(), m * k);
214    debug_assert_eq!(y.len(), n * m);
215    #[cfg(target_os = "macos")]
216    if accel::on() && n * k * m >= 1 << 18 {
217        // Y = X · Wᵀ (row-major).
218        unsafe {
219            accel::cblas_sgemm(
220                101,
221                111,
222                112,
223                n as i32,
224                m as i32,
225                k as i32,
226                1.0,
227                x.as_ptr(),
228                k as i32,
229                w.as_ptr(),
230                k as i32,
231                0.0,
232                y.as_mut_ptr(),
233                m as i32,
234            );
235        }
236        return;
237    }
238    let nb = n.div_ceil(GEMM_BLOCK);
239    let block = |r0: usize, r1: usize, y: &mut [f32]| {
240        for o in 0..m {
241            let wr = &w[o * k..(o + 1) * k];
242            for i in r0..r1 {
243                y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
244            }
245        }
246    };
247    match pool {
248        Some(p) if nb > 1 => {
249            let yp = SendMut(y.as_mut_ptr());
250            p.run(&|widx, nw| {
251                for bi in (widx..nb).step_by(nw) {
252                    let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
253                    // SAFETY: blocks are disjoint row ranges of y.
254                    let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
255                    block(r0, r1, ys);
256                }
257            });
258        }
259        _ => {
260            for bi in 0..nb {
261                let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
262                block(r0, r1, &mut y[r0 * m..r1 * m]);
263            }
264        }
265    }
266}
267
268/// dX += dY[n,m] · W[m,k] — parallel over row blocks (disjoint dX rows).
269pub fn gemm_dx(
270    dy: &[f32],
271    w: &[f32],
272    dx: &mut [f32],
273    n: usize,
274    k: usize,
275    m: usize,
276    pool: Option<&Pool>,
277) {
278    debug_assert_eq!(dy.len(), n * m);
279    debug_assert_eq!(w.len(), m * k);
280    debug_assert_eq!(dx.len(), n * k);
281    #[cfg(target_os = "macos")]
282    if accel::on() && n * k * m >= 1 << 18 {
283        // dX += dY · W (row-major, beta = 1 accumulates).
284        unsafe {
285            accel::cblas_sgemm(
286                101,
287                111,
288                111,
289                n as i32,
290                k as i32,
291                m as i32,
292                1.0,
293                dy.as_ptr(),
294                m as i32,
295                w.as_ptr(),
296                k as i32,
297                1.0,
298                dx.as_mut_ptr(),
299                k as i32,
300            );
301        }
302        return;
303    }
304    let nb = n.div_ceil(GEMM_BLOCK);
305    let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
306        for o in 0..m {
307            let wr = &w[o * k..(o + 1) * k];
308            for i in r0..r1 {
309                let g = dy[i * m + o];
310                if g != 0.0 {
311                    crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
312                }
313            }
314        }
315    };
316    match pool {
317        Some(p) if nb > 1 => {
318            let dxp = SendMut(dx.as_mut_ptr());
319            p.run(&|widx, nw| {
320                for bi in (widx..nb).step_by(nw) {
321                    let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
322                    // SAFETY: blocks are disjoint row ranges of dx.
323                    let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
324                    block(r0, r1, dxs);
325                }
326            });
327        }
328        _ => {
329            for bi in 0..nb {
330                let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
331                block(r0, r1, &mut dx[r0 * k..r1 * k]);
332            }
333        }
334    }
335}
336
337/// dW += dYᵀ · X — parallel over dW ROW ranges (each worker owns a
338/// disjoint slice of output neurons; X is shared read-only).
339pub fn gemm_dw(
340    dy: &[f32],
341    x: &[f32],
342    dw: &mut [f32],
343    n: usize,
344    k: usize,
345    m: usize,
346    pool: Option<&Pool>,
347) {
348    debug_assert_eq!(dy.len(), n * m);
349    debug_assert_eq!(x.len(), n * k);
350    debug_assert_eq!(dw.len(), m * k);
351    #[cfg(target_os = "macos")]
352    if accel::on() && n * k * m >= 1 << 18 {
353        // dW += dYᵀ · X (row-major, beta = 1 accumulates).
354        unsafe {
355            accel::cblas_sgemm(
356                101,
357                112,
358                111,
359                m as i32,
360                k as i32,
361                n as i32,
362                1.0,
363                dy.as_ptr(),
364                m as i32,
365                x.as_ptr(),
366                k as i32,
367                1.0,
368                dw.as_mut_ptr(),
369                k as i32,
370            );
371        }
372        return;
373    }
374    let range = |o0: usize, o1: usize, dws: &mut [f32]| {
375        // i-blocked so the X block stays in cache across the o loop.
376        let nb = n.div_ceil(GEMM_BLOCK);
377        for bi in 0..nb {
378            let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
379            for o in o0..o1 {
380                let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
381                for i in r0..r1 {
382                    let g = dy[i * m + o];
383                    if g != 0.0 {
384                        crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
385                    }
386                }
387            }
388        }
389    };
390    match pool {
391        Some(p) if m >= 8 => {
392            let dwp = SendMut(dw.as_mut_ptr());
393            p.run(&|widx, nw| {
394                let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
395                if o0 < o1 {
396                    // SAFETY: workers own disjoint dW row ranges.
397                    let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
398                    range(o0, o1, dws);
399                }
400            });
401        }
402        _ => range(0, m, dw),
403    }
404}
405
406// ───────────────────────────── SiLU ─────────────────────────────
407
408#[inline]
409pub fn silu<F: Fp>(x: F) -> F {
410    x / (F::ONE + (-x).exp())
411}
412
413/// d silu / dx = σ(x)·(1 + x·(1−σ(x))).
414#[inline]
415pub fn silu_bwd<F: Fp>(x: F) -> F {
416    let s = F::ONE / (F::ONE + (-x).exp());
417    s * (F::ONE + x * (F::ONE - s))
418}
419
420// ──────────────────────────── RMSNorm ────────────────────────────
421
422/// RMSNorm over `n` rows of width `d = w.len()`:
423/// Qwen style y = x̂·w, Gemma style y = x̂·(1+w), x̂ = x/√(mean x² + eps).
424/// Sum-of-squares accumulates in f64 (runtime discipline). `inv_out`
425/// stores the per-row 1/rms for the backward.
426pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
427    let d = w.len();
428    let n = x.len() / d;
429    for r in 0..n {
430        let xr = &x[r * d..(r + 1) * d];
431        let mut ss = 0f64;
432        for v in xr {
433            ss += v.f64() * v.f64();
434        }
435        let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
436        inv_out[r] = inv;
437        let yr = &mut y[r * d..(r + 1) * d];
438        for j in 0..d {
439            let weff = if gemma { F::ONE + w[j] } else { w[j] };
440            yr[j] = xr[j] * inv * weff;
441        }
442    }
443}
444
445/// RMSNorm backward: through-grad into `dx` (+=) and, when the gain is
446/// trainable, gain grad into `dw` (+=). `inv` is the saved 1/rms.
447///
448/// dx_j = inv·w_j·dy_j − x_j·inv³/d · Σ_i dy_i·w_i·x_i
449/// dw_j += dy_j·x_j·inv (identical for both styles: ∂y/∂w = x̂).
450pub fn rmsnorm_bwd<F: Fp>(
451    x: &[F],
452    w: &[F],
453    inv: &[F],
454    dy: &[F],
455    gemma: bool,
456    dx: &mut [F],
457    mut dw: Option<&mut [F]>,
458) {
459    let d = w.len();
460    let n = x.len() / d;
461    for r in 0..n {
462        let xr = &x[r * d..(r + 1) * d];
463        let dyr = &dy[r * d..(r + 1) * d];
464        let iv = inv[r];
465        let mut s = 0f64;
466        for j in 0..d {
467            let weff = if gemma { F::ONE + w[j] } else { w[j] };
468            s += (dyr[j] * weff * xr[j]).f64();
469        }
470        let coef = F::fromf(s / d as f64) * iv * iv * iv;
471        let dxr = &mut dx[r * d..(r + 1) * d];
472        for j in 0..d {
473            let weff = if gemma { F::ONE + w[j] } else { w[j] };
474            dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
475        }
476        if let Some(dwv) = dw.as_deref_mut() {
477            for j in 0..d {
478                dwv[j] += dyr[j] * xr[j] * iv;
479            }
480        }
481    }
482}
483
484// ───────────────────────────── RoPE ─────────────────────────────
485
486/// Rotate the first `2·inv_freq.len()` dims of one head vector in place
487/// (half-split pairing, same convention as `attention::rope_rotate`).
488pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
489    let half = inv_freq.len();
490    for (i, &freq) in inv_freq.iter().enumerate() {
491        let angle = position as f64 * freq;
492        let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
493        let x0 = x[i];
494        let x1 = x[i + half];
495        x[i] = x0 * cos - x1 * sin;
496        x[i + half] = x0 * sin + x1 * cos;
497    }
498}
499
500/// RoPE through-grad: a rotation's Jacobian is the rotation itself, so
501/// dL/dx = R(−θ)·dL/dy — the inverse rotation, in place.
502pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
503    let half = inv_freq.len();
504    for (i, &freq) in inv_freq.iter().enumerate() {
505        let angle = position as f64 * freq;
506        let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
507        let g0 = dy[i];
508        let g1 = dy[i + half];
509        dy[i] = g0 * cos + g1 * sin;
510        dy[i + half] = -g0 * sin + g1 * cos;
511    }
512}
513
514// ─────────────────────── landmark segment means ───────────────────────
515
516/// Contiguous segment means (Nyströmformer landmark recipe), the same
517/// integer split (i·t)/m as `nystrom::seg_means` and the torch probes.
518pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
519    for i in 0..m {
520        let (lo, hi) = (i * t / m, (i + 1) * t / m);
521        let or = &mut out[i * d..(i + 1) * d];
522        for v in or.iter_mut() {
523            *v = F::ZERO;
524        }
525        for j in lo..hi {
526            for c in 0..d {
527                or[c] += x[j * d + c];
528            }
529        }
530        let inv = F::fromf(1.0 / (hi - lo) as f64);
531        for v in or.iter_mut() {
532            *v *= inv;
533        }
534    }
535}
536
537/// Segment-mean backward: the mean is linear, so the landmark grad
538/// scatters back uniformly over its segment (dx += dl/seg_len).
539pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
540    for i in 0..m {
541        let (lo, hi) = (i * t / m, (i + 1) * t / m);
542        let inv = F::fromf(1.0 / (hi - lo) as f64);
543        let dlr = &dl[i * d..(i + 1) * d];
544        for j in lo..hi {
545            for c in 0..d {
546                dx[j * d + c] += dlr[c] * inv;
547            }
548        }
549    }
550}
551
552// ────────────────── exact causal softmax attention ──────────────────
553
554/// Exact per-head causal attention: out[t] = softmax(q_t·Kᵀ/√d)·V over
555/// j ≤ t. `q`,`k` are `[t,d]`, `v` is `[t,dv]`, `out` is `[t,dv]`.
556#[allow(clippy::needless_range_loop)] // row[j] pairs with k-row j — indices are the clearer form
557pub fn attn_head_fwd<F: Fp>(
558    q: &[F],
559    k: &[F],
560    v: &[F],
561    t: usize,
562    d: usize,
563    dv: usize,
564    out: &mut [F],
565) {
566    let scale = F::fromf(1.0 / (d as f64).sqrt());
567    let mut row = vec![F::ZERO; t];
568    for ti in 0..t {
569        let qr = &q[ti * d..(ti + 1) * d];
570        let mut mx = F::fromf(f64::NEG_INFINITY);
571        for j in 0..=ti {
572            let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
573            row[j] = s;
574            mx = mx.maxf(s);
575        }
576        let mut den = F::ZERO;
577        for j in 0..=ti {
578            row[j] = (row[j] - mx).exp();
579            den += row[j];
580        }
581        let or = &mut out[ti * dv..(ti + 1) * dv];
582        for o in or.iter_mut() {
583            *o = F::ZERO;
584        }
585        for j in 0..=ti {
586            let p = row[j] / den;
587            for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
588                *o += p * *vv;
589            }
590        }
591    }
592}
593
594/// Exact-attention backward (probs recomputed row by row — the trainer
595/// is layer-checkpointed, nothing is stored). Standard softmax chain:
596/// ds_j = p_j·(dp_j − Σ p·dp), dp_j = dout·v_j.
597#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
598pub fn attn_head_bwd<F: Fp>(
599    q: &[F],
600    k: &[F],
601    v: &[F],
602    dout: &[F],
603    t: usize,
604    d: usize,
605    dv: usize,
606    dq: &mut [F],
607    dk: &mut [F],
608    dvv: &mut [F],
609) {
610    let scale = F::fromf(1.0 / (d as f64).sqrt());
611    let mut row = vec![F::ZERO; t];
612    for ti in 0..t {
613        let qr = &q[ti * d..(ti + 1) * d];
614        let mut mx = F::fromf(f64::NEG_INFINITY);
615        for j in 0..=ti {
616            let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
617            row[j] = s;
618            mx = mx.maxf(s);
619        }
620        let mut den = F::ZERO;
621        for j in 0..=ti {
622            row[j] = (row[j] - mx).exp();
623            den += row[j];
624        }
625        let dor = &dout[ti * dv..(ti + 1) * dv];
626        // dp_j and the softmax dot Σ p·dp in one pass.
627        let mut pdp = F::ZERO;
628        let mut dp = vec![F::ZERO; ti + 1];
629        for j in 0..=ti {
630            let p = row[j] / den;
631            row[j] = p; // row now holds probabilities
632            dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
633            pdp += p * dp[j];
634        }
635        let dqr = &mut dq[ti * d..(ti + 1) * d];
636        for j in 0..=ti {
637            let p = row[j];
638            // dV
639            for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
640                *dvo += p * *o;
641            }
642            // d logits → dq, dk
643            let ds = p * (dp[j] - pdp) * scale;
644            let kr = &k[j * d..(j + 1) * d];
645            for c in 0..d {
646                dqr[c] += ds * kr[c];
647            }
648            let dkr = &mut dk[j * d..(j + 1) * d];
649            for c in 0..d {
650                dkr[c] += ds * qr[c];
651            }
652        }
653    }
654}
655
656// ─────────────────── Nyström joint attention (matrix form) ───────────────────
657
658/// Nyström joint kernel geometry — matches `nystrom::O1Cfg` semantics.
659#[derive(Clone, Copy, Debug)]
660pub struct NysCfg {
661    /// Landmark budget (m_eff = clamp(t_prefill/8, 4, m), as sealed).
662    pub m: usize,
663    /// Exact-window width.
664    pub w: usize,
665    /// Permanent exact sink keys.
666    pub sink: usize,
667    /// Tokens of each training window that run the EXACT prompt pass
668    /// before the O(1) seal. `None` = half the window.
669    ///
670    /// This is not a free knob: it is the train/serve contract. The
671    /// runtime (`nystrom::NystromState::prefill`) freezes landmarks from
672    /// the PROMPT ONLY and then serves every later position off that
673    /// frozen skeleton, so a trainer that seals from the whole window is
674    /// optimizing a model that is never served. The default mirrors
675    /// `cortiq ppl --o1`'s own default (`o1_prefill = len/2`), which is
676    /// the discipline every published o1 quality number was measured
677    /// with — keep them equal or the polish and the gate disagree.
678    pub prefill: Option<usize>,
679}
680
681impl NysCfg {
682    /// Prompt length for a `t`-token window: the seal point. Clamped to
683    /// [1, t] so a caller cannot ask for a prefix longer than the window.
684    #[inline]
685    pub fn prefill_len(&self, t: usize) -> usize {
686        self.prefill.unwrap_or(t / 2).clamp(1, t)
687    }
688}
689
690/// Joint-denominator floor (mirrors the reference probe / runtime).
691const NYS_DEN_EPS: f64 = 1e-30;
692
693/// Everything the forward computes that the backward reuses. All T×T
694/// buffers — call this with F = f64 (raw exp of real logits overflows
695/// f32; the certified CPU probe ran the skeleton in f64).
696struct NysGraph<F: Fp> {
697    m_eff: usize,
698    /// Seal point: rows < tp are exact (the prompt pass), rows ≥ tp are
699    /// served off the frozen skeleton.
700    tp: usize,
701    q_l: Vec<F>,
702    k_l: Vec<F>,
703    mu: Vec<F>,
704    fu: Vec<F>,
705    e: Vec<F>,
706    fumu: Vec<F>,
707    /// Per row: did the AGGREGATE far denominator survive the guard?
708    /// False ⇒ this row's far field is dropped entirely (weights and
709    /// gradient alike) — see `nys_graph`.
710    ///
711    /// This replaces the raw skeleton estimate `a`, which the backward
712    /// used to carry ONLY to evaluate the per-(t,j) clamp's `[a>0]`
713    /// subgradient mask. The shipped operator gates per ROW, so the mask
714    /// — and the whole T×T buffer behind it — is gone.
715    far_keep: Vec<bool>,
716    /// Final joint weights: exp(lg−c) on near, a·e^{−c} on a kept far
717    /// field, 0 on a dropped one.
718    wmat: Vec<F>,
719    c_row: Vec<F>,
720    den: Vec<F>,
721}
722
723/// Is this row served exactly? Positions before the seal are the
724/// runtime's PROMPT PASS: `Pipeline::nll_ids_o1` runs them through
725/// ordinary full-KV attention and only then calls `o1_seal`, so the
726/// skeleton must not touch them.
727#[inline]
728fn nys_exact_row(ti: usize, tp: usize) -> bool {
729    ti < tp
730}
731
732/// near mask (spec §5b): (t−j < W) OR (j < sink), causal.
733#[inline]
734fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
735    ti - j < w || j < sink
736}
737
738/// Build the joint weight matrix (docs/RUST_FCD.md §2.2). M comes from
739/// the RUNTIME's ridge pseudo-inverse and is CONSTANT in backward.
740/// `mu_override` freezes M explicitly (gradcheck of the constant-M
741/// convention); None recomputes it from the landmarks.
742fn nys_graph<F: Fp>(
743    q: &[F],
744    k: &[F],
745    t: usize,
746    d: usize,
747    cfg: &NysCfg,
748    mu_override: Option<&[F]>,
749) -> NysGraph<F> {
750    let scale = 1.0 / (d as f64).sqrt();
751    let fscale = F::fromf(scale);
752    // Landmarks are sealed from the PROMPT PREFIX, exactly as
753    // `NystromState::prefill` does — including m_eff, which the runtime
754    // derives from the prompt length it sealed at, not from how far the
755    // sequence later runs.
756    let tp = cfg.prefill_len(t);
757    let m_eff = (tp / 8).clamp(4, cfg.m);
758    let mut q_l = vec![F::ZERO; m_eff * d];
759    let mut k_l = vec![F::ZERO; m_eff * d];
760    seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
761    seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
762
763    // Au and its ridge pinv in f64 — one m×m solve, constant in backward.
764    let mut au = vec![0f64; m_eff * m_eff];
765    for i in 0..m_eff {
766        for j in 0..m_eff {
767            let mut s = 0f64;
768            for c in 0..d {
769                s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
770            }
771            au[i * m_eff + j] = (s * scale).exp();
772        }
773    }
774    let mu: Vec<F> = match mu_override {
775        Some(m) => m.to_vec(),
776        None => crate::nystrom::ridge_pinv(&au, m_eff)
777            .iter()
778            .map(|&x| F::fromf(x))
779            .collect(),
780    };
781
782    // Landmark score factors.
783    let mut fu = vec![F::ZERO; t * m_eff];
784    for ti in 0..t {
785        for i in 0..m_eff {
786            fu[ti * m_eff + i] =
787                (dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
788        }
789    }
790    let mut e = vec![F::ZERO; m_eff * t];
791    for i in 0..m_eff {
792        for j in 0..t {
793            e[i * t + j] = (dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
794        }
795    }
796    let mut fumu = vec![F::ZERO; t * m_eff];
797    matmul_nt(
798        &fu,
799        // mu is [m,m] row-major; matmul_nt wants W[m_out, k] rows = muᵀ
800        // columns… avoid transposition juggling: do it directly.
801        &transpose(&mu, m_eff, m_eff),
802        &mut fumu,
803        t,
804        m_eff,
805        m_eff,
806    );
807    // a[t,j] = Σ_i fumu[t,i]·e[i,j] — e is [m,t], so column j of e is
808    // strided; loop with accumulation over i keeps rows contiguous.
809    // Rows before the seal are exact, so their skeleton is never
810    // evaluated (and stays ZERO — the backward relies on that).
811    let mut a = vec![F::ZERO; t * t];
812    for ti in tp..t {
813        let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
814        let ar = &mut a[ti * t..ti * t + ti + 1]; // causal: j ≤ ti
815        for (i, &f) in fr.iter().enumerate() {
816            let er = &e[i * t..i * t + ti + 1];
817            for (av, ev) in ar.iter_mut().zip(er) {
818                *av += f * *ev;
819            }
820        }
821    }
822
823    // Joint weights with the per-row shift c (constant in backward —
824    // it multiplies numerator and denominator identically).
825    let mut wmat = vec![F::ZERO; t * t];
826    let mut c_row = vec![F::ZERO; t];
827    let mut den = vec![F::ZERO; t];
828    let mut far_keep = vec![false; t];
829    let mut lg_row = vec![F::ZERO; t];
830    for ti in 0..t {
831        let qr = &q[ti * d..(ti + 1) * d];
832        let mut c = F::fromf(f64::NEG_INFINITY);
833        for j in 0..=ti {
834            let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
835            lg_row[j] = s;
836            c = c.maxf(s);
837        }
838        c_row[ti] = c;
839        let emc = (-c).exp();
840        let exact_row = nys_exact_row(ti, tp);
841
842        // AGGREGATE guard (`O1Rect::Aggregate`, the shipped default).
843        // The runtime never materializes a per-(t,j) weight — it only
844        // ever holds the far field already contracted into accumulators,
845        // so the ONLY quantity it can test is the row's total far
846        // denominator:
847        //   far_den = Σ_b u[b]·ẑ[b] = e^{−c_all}·Σ_{j far} a[t,j]
848        // (`nystrom.rs::step`). e^{−c_all} > 0, so the runtime's
849        // `far_den >= 0` predicate is exactly `Σ_{j far} a[t,j] >= 0`
850        // here. A row that fails drops its far field ENTIRELY; a row
851        // that passes keeps every far weight RAW — negative per-key mass
852        // included. That is not an oversight: clamping per key (what the
853        // torch probe did, and what this trainer used to do) is a
854        // strictly different, measurably worse operator (`O1Rect::Fm`,
855        // ×1.414 vs ×1.296) that no streaming kernel can execute.
856        let mut far_sum = F::ZERO;
857        if !exact_row {
858            for j in 0..=ti {
859                if !nys_near(ti, j, cfg.w, cfg.sink) {
860                    far_sum += a[ti * t + j];
861                }
862            }
863        }
864        let keep = !exact_row && far_sum.f64() >= 0.0;
865        far_keep[ti] = keep;
866
867        let wr = &mut wmat[ti * t..(ti + 1) * t];
868        let mut dsum = F::ZERO;
869        for j in 0..=ti {
870            let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
871                (lg_row[j] - c).exp()
872            } else if keep {
873                a[ti * t + j] * emc
874            } else {
875                F::ZERO
876            };
877            wr[j] = wv;
878            dsum += wv;
879        }
880        den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
881    }
882    NysGraph {
883        m_eff,
884        tp,
885        q_l,
886        k_l,
887        mu,
888        fu,
889        e,
890        fumu,
891        far_keep,
892        wmat,
893        c_row,
894        den,
895    }
896}
897
898fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
899    let mut out = vec![F::ZERO; rows * cols];
900    for r in 0..rows {
901        for c in 0..cols {
902            out[c * rows + r] = x[r * cols + c];
903        }
904    }
905    out
906}
907
908/// Nyström joint forward (teacher-forced matrix form). Falls back to
909/// exact attention for short windows (runtime guard: t ≤ W+sink+8).
910#[allow(clippy::too_many_arguments)]
911pub fn nystrom_head_fwd<F: Fp>(
912    q: &[F],
913    k: &[F],
914    v: &[F],
915    t: usize,
916    d: usize,
917    dv: usize,
918    cfg: &NysCfg,
919    out: &mut [F],
920) {
921    if nys_degenerate(t, cfg) {
922        attn_head_fwd(q, k, v, t, d, dv, out);
923        return;
924    }
925    nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
926}
927
928/// Does the runtime skip the skeleton for this window entirely?
929///
930/// `NystromState::prefill` decides `exact_only` from the PROMPT length
931/// (`t <= w + sink + EXACT_SLACK`), not from how long the sequence
932/// eventually grows: a short prompt seals no skeleton, so every later
933/// key stays a permanent exact key. Keying this off the full window `t`
934/// (as this did before) would build a skeleton for windows the runtime
935/// serves exactly.
936#[inline]
937fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
938    cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
939}
940
941/// The landmark mixing matrix M of this (q, k) — gradcheck hook for
942/// freezing M across finite-difference perturbations.
943#[doc(hidden)]
944pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
945    nys_graph(q, k, t, d, cfg, None).mu
946}
947
948/// Forward with an explicitly frozen M (gradcheck of the constant-M
949/// convention; the trainer always passes through `nystrom_head_fwd`).
950#[doc(hidden)]
951#[allow(clippy::too_many_arguments)]
952pub fn nystrom_head_fwd_mu<F: Fp>(
953    q: &[F],
954    k: &[F],
955    v: &[F],
956    t: usize,
957    d: usize,
958    dv: usize,
959    cfg: &NysCfg,
960    mu_override: Option<&[F]>,
961    out: &mut [F],
962) {
963    let g = nys_graph(q, k, t, d, cfg, mu_override);
964    for ti in 0..t {
965        let wr = &g.wmat[ti * t..(ti + 1) * t];
966        let den = g.den[ti];
967        let or = &mut out[ti * dv..(ti + 1) * dv];
968        for o in or.iter_mut() {
969            *o = F::ZERO;
970        }
971        for j in 0..=ti {
972            let p = wr[j] / den;
973            if p.f64() != 0.0 {
974                for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
975                    *o += p * *vv;
976                }
977            }
978        }
979    }
980}
981
982/// Nyström joint backward: dq/dk/dv (+=) with M constant. The whole
983/// graph is recomputed (layer checkpointing) — see docs/RUST_FCD.md
984/// §2.2 for the derivation. Chain summary:
985///   dw[t,j]   = dout_t·(v_j − out_t)/den_t
986///   near:  dlg = dw·w                     → dq, dk (±scale)
987///   far:   da  = dw·e^{−c}·[row kept]     → dFu, dE through the two
988///          matmuls with M constant        → dq, dk, dQ̃, dK̃
989///   landmarks: segment-mean scatter back into the PREFIX rows of
990///          dq, dk (the seal only saw those).
991/// The far gate is per ROW (the aggregate guard), not per key.
992#[allow(clippy::too_many_arguments)]
993pub fn nystrom_head_bwd<F: Fp>(
994    q: &[F],
995    k: &[F],
996    v: &[F],
997    dout: &[F],
998    t: usize,
999    d: usize,
1000    dv: usize,
1001    cfg: &NysCfg,
1002    dq: &mut [F],
1003    dk: &mut [F],
1004    dvv: &mut [F],
1005) {
1006    if nys_degenerate(t, cfg) {
1007        attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
1008        return;
1009    }
1010    nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
1011}
1012
1013/// Backward with an explicitly frozen M (see `nystrom_head_fwd_mu`).
1014#[doc(hidden)]
1015#[allow(clippy::too_many_arguments)]
1016pub fn nystrom_head_bwd_mu<F: Fp>(
1017    q: &[F],
1018    k: &[F],
1019    v: &[F],
1020    dout: &[F],
1021    t: usize,
1022    d: usize,
1023    dv: usize,
1024    cfg: &NysCfg,
1025    mu_override: Option<&[F]>,
1026    dq: &mut [F],
1027    dk: &mut [F],
1028    dvv: &mut [F],
1029) {
1030    let scale = F::fromf(1.0 / (d as f64).sqrt());
1031    let g = nys_graph(q, k, t, d, cfg, mu_override);
1032    let m_eff = g.m_eff;
1033
1034    // Recompute out rows (needed inside dw), then dv and dw in one pass.
1035    let mut dwmat = vec![F::ZERO; t * t];
1036    let mut out_row = vec![F::ZERO; dv];
1037    for ti in 0..t {
1038        let wr = &g.wmat[ti * t..(ti + 1) * t];
1039        let den = g.den[ti];
1040        for o in out_row.iter_mut() {
1041            *o = F::ZERO;
1042        }
1043        for j in 0..=ti {
1044            let p = wr[j] / den;
1045            if p.f64() != 0.0 {
1046                for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1047                    *o += p * *vv;
1048                }
1049            }
1050        }
1051        let dor = &dout[ti * dv..(ti + 1) * dv];
1052        let dwr = &mut dwmat[ti * t..(ti + 1) * t];
1053        for j in 0..=ti {
1054            // dV: p·dout
1055            let p = wr[j] / den;
1056            if p.f64() != 0.0 {
1057                for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
1058                    *dvo += p * *o;
1059                }
1060            }
1061            // dw = dout·(v_j − out)/den
1062            let mut s = F::ZERO;
1063            for c in 0..dv {
1064                s += dor[c] * (v[j * dv + c] - out_row[c]);
1065            }
1066            dwr[j] = s / den;
1067        }
1068    }
1069
1070    // Near half: dlg = dw·w, straight to dq/dk. A pre-seal row is
1071    // exact, so EVERY causal j is a near key for it.
1072    for ti in 0..t {
1073        let qr = &q[ti * d..(ti + 1) * d];
1074        let dqr_base = ti * d;
1075        for j in 0..=ti {
1076            if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
1077                continue;
1078            }
1079            let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
1080            if dlg.f64() == 0.0 {
1081                continue;
1082            }
1083            let kr = &k[j * d..(j + 1) * d];
1084            for c in 0..d {
1085                dq[dqr_base + c] += dlg * kr[c];
1086            }
1087            let dkr = &mut dk[j * d..(j + 1) * d];
1088            for c in 0..d {
1089                dkr[c] += dlg * qr[c];
1090            }
1091        }
1092    }
1093
1094    // Far half: da gated by the AGGREGATE guard, then back through the
1095    // skeleton.
1096    //
1097    // Gradient of the guard: the far field enters as the piecewise
1098    // function  far(row) = [Σ_far a ≥ 0] · a·e^{−c}, i.e. a per-ROW
1099    // switch, not a per-key clamp. On the kept branch it is the identity
1100    // in a — so the far grad flows RAW, with no [a>0] mask (that mask
1101    // belonged to the per-key clamp this replaces and would now zero the
1102    // gradient of exactly the negative-mass keys the shipped operator
1103    // keeps). On the dropped branch the far field is identically 0 in a
1104    // neighbourhood of a, so its gradient is 0. The switching boundary
1105    // itself (Σ_far a = 0) is measure-zero and, like any subgradient
1106    // convention, is not differentiated — at that point far ≡ 0 anyway,
1107    // so the two branches agree in value and only the derivative jumps.
1108    let mut da = vec![F::ZERO; t * t];
1109    for ti in g.tp..t {
1110        if !g.far_keep[ti] {
1111            continue; // dropped row: zero gradient through its far field
1112        }
1113        let emc = (-g.c_row[ti]).exp();
1114        for j in 0..=ti {
1115            if nys_near(ti, j, cfg.w, cfg.sink) {
1116                continue;
1117            }
1118            da[ti * t + j] = dwmat[ti * t + j] * emc;
1119        }
1120    }
1121    // dFuMu[t,i] = Σ_j da[t,j]·e[i,j]
1122    let mut dfumu = vec![F::ZERO; t * m_eff];
1123    for ti in 0..t {
1124        let dar = &da[ti * t..ti * t + ti + 1];
1125        let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
1126        for (i, df) in dfr.iter_mut().enumerate() {
1127            let er = &g.e[i * t..i * t + ti + 1];
1128            let mut s = F::ZERO;
1129            for (av, ev) in dar.iter().zip(er) {
1130                s += *av * *ev;
1131            }
1132            *df = s;
1133        }
1134    }
1135    // dFu = dFuMu·Muᵀ  (M constant)
1136    let mut dfu = vec![F::ZERO; t * m_eff];
1137    matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
1138    // dE[i,j] = Σ_t fumu[t,i]·da[t,j]
1139    let mut de = vec![F::ZERO; m_eff * t];
1140    for ti in 0..t {
1141        let dar = &da[ti * t..ti * t + ti + 1];
1142        let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
1143        for (i, &f) in fr.iter().enumerate() {
1144            if f.f64() == 0.0 {
1145                continue;
1146            }
1147            let der = &mut de[i * t..i * t + ti + 1];
1148            for (dev, av) in der.iter_mut().zip(dar) {
1149                *dev += f * *av;
1150            }
1151        }
1152    }
1153    // Chain through the two exponentials into dq/dk and the landmark
1154    // grads dQ̃/dK̃.
1155    let mut dq_l = vec![F::ZERO; m_eff * d];
1156    let mut dk_l = vec![F::ZERO; m_eff * d];
1157    for ti in 0..t {
1158        let qr = &q[ti * d..(ti + 1) * d];
1159        for i in 0..m_eff {
1160            let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
1161            if dlg.f64() == 0.0 {
1162                continue;
1163            }
1164            let klr = &g.k_l[i * d..(i + 1) * d];
1165            for c in 0..d {
1166                dq[ti * d + c] += dlg * klr[c];
1167            }
1168            let dklr = &mut dk_l[i * d..(i + 1) * d];
1169            for c in 0..d {
1170                dklr[c] += dlg * qr[c];
1171            }
1172        }
1173    }
1174    for i in 0..m_eff {
1175        let qlr = &g.q_l[i * d..(i + 1) * d];
1176        for j in 0..t {
1177            let dlg = de[i * t + j] * g.e[i * t + j] * scale;
1178            if dlg.f64() == 0.0 {
1179                continue;
1180            }
1181            let kr = &k[j * d..(j + 1) * d];
1182            let dqlr = &mut dq_l[i * d..(i + 1) * d];
1183            for c in 0..d {
1184                dqlr[c] += dlg * kr[c];
1185            }
1186            for c in 0..d {
1187                dk[j * d + c] += dlg * qlr[c];
1188            }
1189        }
1190    }
1191    // Landmarks were sealed from the prompt prefix, so their gradient
1192    // scatters back into those rows ONLY — a post-seal token cannot move
1193    // a landmark it never contributed to.
1194    let tp = g.tp;
1195    seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
1196    seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
1197}
1198
1199// ───────────────────────────── losses ─────────────────────────────
1200
1201/// One position of the polish loss (docs/RUST_FCD.md §2.4):
1202/// L = (1−klw)·CE(student, target) + klw·KL(teacher‖student), both
1203/// per-position; `inv_n` = 1/(B·T) folds the batch mean into dlogits.
1204/// Returns (ce, kl) UNWEIGHTED for logging; dlogits gets the combined
1205/// gradient (+= is NOT used — each position owns its slice).
1206pub fn ce_kl_position<F: Fp>(
1207    s_logits: &[F],
1208    t_logits: &[F],
1209    target: usize,
1210    kl_w: f64,
1211    inv_n: f64,
1212    dlogits: &mut [F],
1213) -> (f64, f64) {
1214    let vsz = s_logits.len();
1215    debug_assert_eq!(t_logits.len(), vsz);
1216    // Student log-softmax in f64.
1217    let mut smax = f64::NEG_INFINITY;
1218    let mut tmax = f64::NEG_INFINITY;
1219    for i in 0..vsz {
1220        smax = smax.max(s_logits[i].f64());
1221        tmax = tmax.max(t_logits[i].f64());
1222    }
1223    let mut ssum = 0f64;
1224    let mut tsum = 0f64;
1225    for i in 0..vsz {
1226        ssum += (s_logits[i].f64() - smax).exp();
1227        tsum += (t_logits[i].f64() - tmax).exp();
1228    }
1229    let slz = smax + ssum.ln();
1230    let tlz = tmax + tsum.ln();
1231    let ce = slz - s_logits[target].f64();
1232    let mut kl = 0f64;
1233    for i in 0..vsz {
1234        let ls = s_logits[i].f64() - slz;
1235        let lt = t_logits[i].f64() - tlz;
1236        let pt = lt.exp();
1237        let ps = ls.exp();
1238        if pt > 0.0 {
1239            kl += pt * (lt - ls);
1240        }
1241        let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
1242        if i == target {
1243            gd -= 1.0 - kl_w;
1244        }
1245        dlogits[i] = F::fromf(gd * inv_n);
1246    }
1247    (ce, kl)
1248}
1249
1250// ─────────────────── GatedDeltaNet through-backward ───────────────────
1251//
1252// Faithful BPTT through `linear_core::gdn_step` semantics (docs/
1253// RUST_FCD.md §3): depthwise causal conv + SiLU, per-group l2-normalized
1254// q/k, gates g = exp(−exp(A_log)·softplus(a+dt_bias)) and β = σ(b), the
1255// delta-rule recurrence S ← g·S; kv = Sᵀk̂; S += k̂⊗β(v−kv); o = Sᵀq̂, and
1256// the gated per-head RMSNorm output x̂·w·silu(z). Through-grad ONLY —
1257// every GDN weight stays frozen (the FCD policy: attention operators
1258// are closed-form/frozen, training touches LN+FFN of converted layers).
1259//
1260// The backward stores the full per-head state history S_0..S_T (one
1261// head at a time: T·dk·dv floats — ~67 MB f64 at Qwen3.5-0.8B geometry,
1262// freed per head). Larger models would switch to segment checkpoints;
1263// the entry points below don't change for that.
1264
1265/// GDN geometry + frozen elementwise weights for the sequence ops.
1266pub struct GdnSeqCfg<'a> {
1267    pub nv: usize,
1268    pub nk: usize,
1269    pub dk: usize,
1270    pub dv: usize,
1271    pub kk: usize,
1272    pub rms_eps: f64,
1273    /// Depthwise conv taps `[c_dim × kk]`, oldest→newest (tap kk−1
1274    /// multiplies the current position) — `GdnWeights::conv1d` layout.
1275    pub conv: &'a [f32],
1276    /// Per-v-head decay parameter A_log `[nv]`.
1277    pub a_log: &'a [f32],
1278    /// Per-v-head dt bias `[nv]`.
1279    pub dt_bias: &'a [f32],
1280    /// Gated-RMSNorm gain `[dv]` (plain x̂·w, norm-before-gate).
1281    pub norm: &'a [f32],
1282}
1283
1284impl GdnSeqCfg<'_> {
1285    pub fn c_dim(&self) -> usize {
1286        2 * self.nk * self.dk + self.nv * self.dv
1287    }
1288}
1289
1290#[inline]
1291fn softplus_f<F: Fp>(x: F) -> F {
1292    // Same threshold as linear_core::softplus — parity matters more
1293    // than elegance (σ(20) ≈ 1 − 2e-9, consistent with the cutoff).
1294    if x.f64() > 20.0 {
1295        x
1296    } else {
1297        F::fromf(x.f64().exp().ln_1p())
1298    }
1299}
1300
1301#[inline]
1302fn sigmoid_f<F: Fp>(x: F) -> F {
1303    F::ONE / (F::ONE + (-x).exp())
1304}
1305
1306/// Depthwise causal conv + SiLU over the whole window: raw `[t, c_dim]`
1307/// → (pre-activation `[t, c_dim]`, cq = silu(pre)). Ring semantics of
1308/// `gdn_step` from a fresh state: positions before 0 are zeros.
1309pub fn gdn_conv_fwd<F: Fp>(
1310    raw: &[F],
1311    t: usize,
1312    c_dim: usize,
1313    kk: usize,
1314    conv: &[f32],
1315    pre: &mut [F],
1316    cq: &mut [F],
1317) {
1318    for ti in 0..t {
1319        for c in 0..c_dim {
1320            let taps = &conv[c * kk..(c + 1) * kk];
1321            let mut acc = F::ZERO;
1322            for (j, &tap) in taps.iter().enumerate() {
1323                // tap j multiplies raw position ti − (kk−1) + j.
1324                let p = ti as isize - (kk as isize - 1) + j as isize;
1325                if p >= 0 {
1326                    acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
1327                }
1328            }
1329            pre[ti * c_dim + c] = acc;
1330            cq[ti * c_dim + c] = silu(acc);
1331        }
1332    }
1333}
1334
1335/// Conv+SiLU backward: dcq → draw (+=). Taps frozen (through-grad only).
1336pub fn gdn_conv_bwd<F: Fp>(
1337    pre: &[F],
1338    t: usize,
1339    c_dim: usize,
1340    kk: usize,
1341    conv: &[f32],
1342    dcq: &[F],
1343    draw: &mut [F],
1344) {
1345    for ti in 0..t {
1346        for c in 0..c_dim {
1347            let g = dcq[ti * c_dim + c];
1348            if g.f64() == 0.0 {
1349                continue;
1350            }
1351            let dp = g * silu_bwd(pre[ti * c_dim + c]);
1352            let taps = &conv[c * kk..(c + 1) * kk];
1353            for (j, &tap) in taps.iter().enumerate() {
1354                let p = ti as isize - (kk as isize - 1) + j as isize;
1355                if p >= 0 {
1356                    draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
1357                }
1358            }
1359        }
1360    }
1361}
1362
1363/// l2-normalization factors of one q/k vector, matching gdn_step:
1364/// invq = 1/(√(Σq²+1e-6)·√dk), invk = 1/√(Σk²+1e-6).
1365#[inline]
1366fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
1367    let mut n2 = F::ZERO;
1368    for v in x {
1369        n2 += *v * *v;
1370    }
1371    let n2e = n2 + F::fromf(1e-6);
1372    let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
1373    (inv, n2e)
1374}
1375
1376/// One GQA group (k-head `ko`, its rep = nv/nk v-heads) forward over
1377/// the window. Writes `out[t, nv·dv]` slices of its v-heads only.
1378#[allow(clippy::too_many_arguments)]
1379pub fn gdn_group_fwd<F: Fp>(
1380    cq: &[F],
1381    z: &[F],
1382    a: &[F],
1383    b: &[F],
1384    t: usize,
1385    cfg: &GdnSeqCfg,
1386    ko: usize,
1387    out: &mut [F],
1388) {
1389    let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1390    let c_dim = cfg.c_dim();
1391    let kd = nk * dk;
1392    let rep = nv / nk;
1393    let vd = nv * dv;
1394    let sqdk = (dk as f64).sqrt();
1395    for hh in 0..rep {
1396        let h = ko * rep + hh;
1397        let ea = F::fromf((cfg.a_log[h] as f64).exp());
1398        let mut s = vec![F::ZERO; dk * dv];
1399        let mut kv = vec![F::ZERO; dv];
1400        let mut o = vec![F::ZERO; dv];
1401        for ti in 0..t {
1402            let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1403            let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1404            let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1405            let (invq, _) = gdn_inv(qrow, sqdk);
1406            let (invk, _) = gdn_inv(krow, 1.0);
1407            let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
1408            let beta = sigmoid_f(b[ti * nv + h]);
1409            // S ← g·S; kv = Sᵀk̂; S += k̂ ⊗ β(v − kv); o = Sᵀq̂.
1410            for x in kv.iter_mut() {
1411                *x = F::ZERO;
1412            }
1413            for di in 0..dk {
1414                let kf = krow[di] * invk;
1415                let row = &mut s[di * dv..(di + 1) * dv];
1416                for dj in 0..dv {
1417                    row[dj] *= g;
1418                    kv[dj] += row[dj] * kf;
1419                }
1420            }
1421            for x in o.iter_mut() {
1422                *x = F::ZERO;
1423            }
1424            for di in 0..dk {
1425                let kf = krow[di] * invk;
1426                let qf = qrow[di] * invq;
1427                let row = &mut s[di * dv..(di + 1) * dv];
1428                for dj in 0..dv {
1429                    row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
1430                    o[dj] += qf * row[dj];
1431                }
1432            }
1433            // Gated per-head RMSNorm: x̂·w·silu(z), norm BEFORE gate.
1434            let mut ss = 0f64;
1435            for v in &o {
1436                ss += v.f64() * v.f64();
1437            }
1438            let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1439            for dj in 0..dv {
1440                let zv = z[ti * vd + h * dv + dj];
1441                out[ti * vd + h * dv + dj] = o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
1442            }
1443        }
1444    }
1445}
1446
1447/// One GQA group backward: BPTT over the window given `dout` rows of
1448/// this group's v-heads. Accumulates (+=) into dcq (q/k channels of
1449/// `ko`, v channels of its v-heads), dz, da, db. All weights frozen.
1450#[allow(clippy::too_many_arguments)]
1451pub fn gdn_group_bwd<F: Fp>(
1452    cq: &[F],
1453    z: &[F],
1454    a: &[F],
1455    b: &[F],
1456    t: usize,
1457    cfg: &GdnSeqCfg,
1458    ko: usize,
1459    dout: &[F],
1460    dcq: &mut [F],
1461    dz: &mut [F],
1462    da: &mut [F],
1463    db: &mut [F],
1464) {
1465    let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1466    let c_dim = cfg.c_dim();
1467    let kd = nk * dk;
1468    let rep = nv / nk;
1469    let vd = nv * dv;
1470    let sqdk = (dk as f64).sqrt();
1471    for hh in 0..rep {
1472        let h = ko * rep + hh;
1473        let ea = F::fromf((cfg.a_log[h] as f64).exp());
1474
1475        // ── replay forward, keeping the state history + per-step scalars ──
1476        let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; // S_0 = 0
1477        let mut kv_hist = vec![F::ZERO; t * dv];
1478        let mut o_hist = vec![F::ZERO; t * dv];
1479        let mut g_v = vec![F::ZERO; t];
1480        let mut beta_v = vec![F::ZERO; t];
1481        let mut sp_arg = vec![F::ZERO; t]; // a + dt_bias (for σ in the chain)
1482        for ti in 0..t {
1483            let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1484            let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1485            let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1486            let (invq, _) = gdn_inv(qrow, sqdk);
1487            let (invk, _) = gdn_inv(krow, 1.0);
1488            let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
1489            let g = (-ea * softplus_f(arg)).exp();
1490            let beta = sigmoid_f(b[ti * nv + h]);
1491            sp_arg[ti] = arg;
1492            g_v[ti] = g;
1493            beta_v[ti] = beta;
1494            let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
1495            let sp = &prev[ti * dk * dv..];
1496            let sn = &mut cur[..dk * dv];
1497            let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
1498            for di in 0..dk {
1499                let kf = krow[di] * invk;
1500                for dj in 0..dv {
1501                    let dec = sp[di * dv + dj] * g;
1502                    sn[di * dv + dj] = dec;
1503                    kvr[dj] += dec * kf;
1504                }
1505            }
1506            let or = &mut o_hist[ti * dv..(ti + 1) * dv];
1507            for di in 0..dk {
1508                let kf = krow[di] * invk;
1509                let qf = qrow[di] * invq;
1510                for dj in 0..dv {
1511                    let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
1512                    sn[di * dv + dj] = sv;
1513                    or[dj] += qf * sv;
1514                }
1515            }
1516        }
1517
1518        // ── reverse sweep ──
1519        let mut ds = vec![F::ZERO; dk * dv];
1520        let mut do_o = vec![F::ZERO; dv];
1521        let mut du = vec![F::ZERO; dv];
1522        let mut dkv = vec![F::ZERO; dv];
1523        let mut dqh = vec![F::ZERO; dk];
1524        let mut dkh = vec![F::ZERO; dk];
1525        for ti in (0..t).rev() {
1526            let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1527            let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1528            let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1529            let (invq, nq2) = gdn_inv(qrow, sqdk);
1530            let (invk, nk2) = gdn_inv(krow, 1.0);
1531            let g = g_v[ti];
1532            let beta = beta_v[ti];
1533            let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
1534            let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
1535            let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
1536            let or = &o_hist[ti * dv..(ti + 1) * dv];
1537
1538            // 1. Gated RMSNorm output: of = (o·inv)·w·silu(z).
1539            let mut ss = 0f64;
1540            for v in or {
1541                ss += v.f64() * v.f64();
1542            }
1543            let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1544            let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
1545            // dz and the effective-gain through-grad in one pass.
1546            let mut sdot = 0f64; // Σ dof·weff·o (for the rms through term)
1547            for dj in 0..dv {
1548                let zv = z[ti * vd + h * dv + dj];
1549                let w = F::fromf(cfg.norm[dj] as f64);
1550                let weff = w * silu(zv);
1551                sdot += (dofr[dj] * weff * or[dj]).f64();
1552                dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
1553            }
1554            let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
1555            for dj in 0..dv {
1556                let zv = z[ti * vd + h * dv + dj];
1557                let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
1558                do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
1559            }
1560
1561            // 2. o = S_tᵀ q̂ → dS += q̂ ⊗ do, dq̂ = S_t·do.
1562            for x in dqh.iter_mut() {
1563                *x = F::ZERO;
1564            }
1565            for di in 0..dk {
1566                let qf = qrow[di] * invq;
1567                let row = &s_t[di * dv..(di + 1) * dv];
1568                let dsr = &mut ds[di * dv..(di + 1) * dv];
1569                let mut acc = F::ZERO;
1570                for dj in 0..dv {
1571                    dsr[dj] += qf * do_o[dj];
1572                    acc += row[dj] * do_o[dj];
1573                }
1574                dqh[di] = acc;
1575            }
1576
1577            // 3. S_t = S_pre + k̂ ⊗ u, u = β(v − kv):
1578            //    du = dSᵀk̂; dk̂ = dS·u; dβ = du·(v−kv); dv = β·du; dkv = −β·du.
1579            for x in du.iter_mut() {
1580                *x = F::ZERO;
1581            }
1582            for x in dkh.iter_mut() {
1583                *x = F::ZERO;
1584            }
1585            for di in 0..dk {
1586                let kf = krow[di] * invk;
1587                let dsr = &ds[di * dv..(di + 1) * dv];
1588                let mut acc = F::ZERO;
1589                for dj in 0..dv {
1590                    du[dj] += dsr[dj] * kf;
1591                    acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
1592                }
1593                dkh[di] = acc;
1594            }
1595            let mut dbeta = F::ZERO;
1596            for dj in 0..dv {
1597                dbeta += du[dj] * (vrow[dj] - kvr[dj]);
1598                // v channels of cq
1599                dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
1600                dkv[dj] = -(beta * du[dj]);
1601            }
1602
1603            // 4. kv = S_preᵀk̂ → dS_pre = dS + k̂ ⊗ dkv; dk̂ += S_pre·dkv.
1604            //    5. S_pre = g·S_{t−1} → dg = ⟨dS_pre, S_{t−1}⟩; carry
1605            //    dS = g·dS_pre to the previous step. (S_pre = g·s_prev
1606            //    is rebuilt on the fly.)
1607            let mut dg = F::ZERO;
1608            for di in 0..dk {
1609                let kf = krow[di] * invk;
1610                let spr = &s_prev[di * dv..(di + 1) * dv];
1611                let dsr = &mut ds[di * dv..(di + 1) * dv];
1612                let mut acc = F::ZERO;
1613                for dj in 0..dv {
1614                    let dspre = dsr[dj] + kf * dkv[dj];
1615                    acc += (spr[dj] * g) * dkv[dj];
1616                    dg += dspre * spr[dj];
1617                    dsr[dj] = g * dspre;
1618                }
1619                dkh[di] += acc;
1620            }
1621
1622            // 6. Gates: g = exp(−e_A·softplus(arg)), β = σ(b).
1623            let sig = sigmoid_f(sp_arg[ti]);
1624            da[ti * nv + h] += dg * g * (-ea) * sig;
1625            db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
1626
1627            // 7. l2-norm through-grads into the shared q/k channels.
1628            //    q̂ = q·invq with invq = 1/(√(Σq²+eps)·√dk):
1629            //    dq = invq·dq̂ − q·(dq̂·q)·invq/(Σq²+eps).
1630            let mut qdot = F::ZERO;
1631            let mut kdot = F::ZERO;
1632            for di in 0..dk {
1633                qdot += dqh[di] * qrow[di];
1634                kdot += dkh[di] * krow[di];
1635            }
1636            for di in 0..dk {
1637                dcq[ti * c_dim + ko * dk + di] += invq * dqh[di] - qrow[di] * qdot * invq / nq2;
1638                dcq[ti * c_dim + kd + ko * dk + di] +=
1639                    invk * dkh[di] - krow[di] * kdot * invk / nk2;
1640            }
1641        }
1642    }
1643}
1644
1645/// Whole-layer GDN sequence forward (serial over groups): raw
1646/// projections → out `[t, nv·dv]`. The trainer parallelizes groups on
1647/// the pool with the same `gdn_group_*` entry points.
1648pub fn gdn_seq_fwd<F: Fp>(
1649    qkv: &[F],
1650    z: &[F],
1651    a: &[F],
1652    b: &[F],
1653    t: usize,
1654    cfg: &GdnSeqCfg,
1655    out: &mut [F],
1656) {
1657    let c_dim = cfg.c_dim();
1658    let mut pre = vec![F::ZERO; t * c_dim];
1659    let mut cq = vec![F::ZERO; t * c_dim];
1660    gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1661    for ko in 0..cfg.nk {
1662        gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
1663    }
1664}
1665
1666/// Whole-layer GDN sequence backward: through-grads (+=) into the four
1667/// projection streams.
1668#[allow(clippy::too_many_arguments)]
1669pub fn gdn_seq_bwd<F: Fp>(
1670    qkv: &[F],
1671    z: &[F],
1672    a: &[F],
1673    b: &[F],
1674    t: usize,
1675    cfg: &GdnSeqCfg,
1676    dout: &[F],
1677    dqkv: &mut [F],
1678    dz: &mut [F],
1679    da: &mut [F],
1680    db: &mut [F],
1681) {
1682    let c_dim = cfg.c_dim();
1683    let mut pre = vec![F::ZERO; t * c_dim];
1684    let mut cq = vec![F::ZERO; t * c_dim];
1685    gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1686    let mut dcq = vec![F::ZERO; t * c_dim];
1687    for ko in 0..cfg.nk {
1688        gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
1689    }
1690    gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
1691}