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