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