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