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