Skip to main content

cortiq_engine/
fcd.rs

1//! FCD polish trainer — native Rust quality-polish for O(1)-converted
2//! models (docs/RUST_FCD.md). Removes the last Python dependency from
3//! the `cortiq convert --o1` pipeline.
4//!
5//! Certified recipe (torch reference `nystrom_fcd_full2_06b.py`,
6//! Qwen3-0.6B 28/28): train ONLY the LN gains + FFN of the converted
7//! layers, loss = (1−0.7)·CE + 0.7·KL(teacher‖student), AdamW 5e-5
8//! (torch defaults), grad clip 1.0, batch 2×512 fresh random windows,
9//! quick deterministic val every 25 steps, restore the best checkpoint.
10//!
11//! Structure:
12//! - the whole model is dequantized to f32 once; teacher and student
13//!   SHARE the frozen set, trainables get separate master copies (the
14//!   KL anchor never drifts);
15//! - layer-level activation checkpointing: the forward keeps only each
16//!   layer's input hidden; the backward re-runs one layer at a time;
17//! - converted layers use the certified matrix form of the Nyström
18//!   joint kernel in f64 (fcd_ops), M constant in backward;
19//! - GDN (GatedDeltaNet) layers of Qwen3.5-class hybrids run frozen in
20//!   BOTH teacher and student, with a true BPTT through-backward
21//!   (fcd_ops::gdn_*) so trainable layers BELOW them still learn;
22//!   vmf_phase linear layers are refused (no backward yet).
23
24use crate::fcd_ops::{self as ops, NysCfg};
25use crate::nystrom::{O1Cfg, O1Layers};
26use crate::pipeline::{DenseFfn, FfnKind, Pipeline};
27use crate::pool::Pool;
28use crate::qtensor::QTensor;
29use crate::sampler::{SamplerConfig, SplitMix64};
30use cortiq_core::{CmfModel, LayerType, NormStyle, TensorDtype};
31use std::sync::Arc;
32
33/// Position-chunk size for the tied-head loss (never materializes the
34/// full [B·T, vocab] logits of teacher AND student together).
35const LM_CHUNK: usize = 32;
36
37/// AdamW hyper-parameters — torch defaults, part of the certified recipe.
38const ADAM_B1: f64 = 0.9;
39const ADAM_B2: f64 = 0.999;
40const ADAM_EPS: f64 = 1e-8;
41const ADAM_WD: f64 = 0.01;
42
43/// Training hyper-parameters (defaults = the certified recipe).
44#[derive(Clone, Debug)]
45pub struct FcdHyper {
46    pub steps: usize,
47    pub lr: f64,
48    pub kl_w: f64,
49    pub eval_every: usize,
50    pub bs: usize,
51    pub seq: usize,
52    pub seed: u64,
53}
54
55impl Default for FcdHyper {
56    fn default() -> Self {
57        Self {
58            steps: 300,
59            lr: 5e-5,
60            kl_w: 0.7,
61            eval_every: 25,
62            bs: 2,
63            seq: 512,
64            seed: 0,
65        }
66    }
67}
68
69/// What the polish measured — written into `provenance.fcd` and
70/// reported by the CLI.
71#[derive(Clone, Debug)]
72pub struct FcdReport {
73    pub converted: Vec<usize>,
74    /// Teacher (exact attention) quick-val ppl — the anchor.
75    pub teacher_ppl: f64,
76    /// Student quick-val ppl BEFORE training (zero-shot o1 damage).
77    pub ppl_start: f64,
78    /// Best quick-val ppl during training (the restored checkpoint).
79    pub ppl_best: f64,
80    pub best_step: usize,
81    /// Final val ppl of the restored checkpoint on the wider window set.
82    pub ppl_final: f64,
83    pub steps_run: usize,
84    pub sec_per_step: f64,
85    /// Per-step (ce, kl) — the training trajectory, unweighted.
86    pub losses: Vec<(f64, f64)>,
87    /// Generation-gate record (None = ppl-only selection).
88    pub gate: Option<GateReport>,
89}
90
91/// What the generation gate saw and decided.
92#[derive(Clone, Debug)]
93pub struct GateReport {
94    /// Zero-shot (step-0) loop scores per prompt — the baseline.
95    pub baseline: Vec<f64>,
96    /// Per eval checkpoint: (step, val ppl, loop scores, passed).
97    pub evals: Vec<(usize, f64, Vec<f64>, bool)>,
98    /// Step whose params were restored (None = identity: the polish
99    /// was rejected, the artifact carries the zero-shot state).
100    pub chosen: Option<usize>,
101}
102
103// ───────────────────── generation gate (claim 13) ─────────────────────
104
105/// Loopiness of a generated id sequence: 1 − unique 4-grams / total.
106/// 0 = no repeated 4-gram; near 1 = a tight loop.
107pub fn loop_score(ids: &[u32]) -> f64 {
108    if ids.len() < 5 {
109        return 0.0;
110    }
111    let grams: std::collections::HashSet<&[u32]> = ids.windows(4).collect();
112    1.0 - grams.len() as f64 / ids.windows(4).count() as f64
113}
114
115/// Generation-gate configuration (Patent 16 draft, claim 13:
116/// checkpoint selection gated on generation-behavior metrics measured
117/// through the SERVED kernel, not on the training objective alone).
118#[derive(Clone, Debug)]
119pub struct GenGateCfg {
120    /// Fixed long-context prompts (token ids), greedy-decoded at every
121    /// eval checkpoint.
122    pub prompts: Vec<Vec<u32>>,
123    pub gen_tokens: usize,
124    /// A checkpoint fails if ANY prompt's loop score exceeds this.
125    pub threshold: f64,
126    /// …or exceeds its zero-shot baseline by more than this.
127    pub baseline_slack: f64,
128}
129
130impl GenGateCfg {
131    /// The standard 3-prompt probe of the torch reference: 400-token
132    /// windows at L/10, L/2, 8L/10 of the val stream, greedy 60.
133    pub fn standard(va: &[u32]) -> Option<Self> {
134        let l = va.len().saturating_sub(500);
135        if l < 400 {
136            return None;
137        }
138        let prompts = [l / 10, l / 2, 8 * l / 10]
139            .iter()
140            .map(|&off| va[off..off + 400].to_vec())
141            .collect();
142        Some(Self {
143            prompts,
144            gen_tokens: 60,
145            threshold: 0.35,
146            baseline_slack: 0.10,
147        })
148    }
149}
150
151/// Gate predicate (Patent 16 draft, claim 13): a checkpoint PASSES iff
152/// no prompt's loop score exceeds `threshold` AND none exceeds its
153/// zero-shot baseline by more than `slack` — boundary values pass.
154pub fn gate_pass(scores: &[f64], baseline: &[f64], threshold: f64, slack: f64) -> bool {
155    scores
156        .iter()
157        .zip(baseline)
158        .all(|(&s, &b)| s <= threshold && s <= b + slack)
159}
160
161/// Checkpoint selection: lowest val ppl AMONG GATE-PASSING checkpoints
162/// (ties → earliest). None = nothing passed → the caller must restore
163/// the zero-shot state (identity polish): the stage must never make
164/// generation worse than conversion alone. (Patent 16 draft, claim 13.)
165pub fn select_checkpoint(
166    evals: &[(usize, f64, Vec<f64>)],
167    baseline: &[f64],
168    threshold: f64,
169    slack: f64,
170) -> Option<usize> {
171    let mut best: Option<usize> = None;
172    for (i, (_, ppl, scores)) in evals.iter().enumerate() {
173        if !gate_pass(scores, baseline, threshold, slack) {
174            continue;
175        }
176        if best.map(|b| *ppl < evals[b].1).unwrap_or(true) {
177            best = Some(i);
178        }
179    }
180    best
181}
182
183// ───────────────────────── model container ─────────────────────────
184
185
186/// Wall-clock phase counters for the bake profiler (CMF_BAKE_PROF=1
187/// prints them). Written from worker threads, so atomics; nanoseconds.
188pub mod prof {
189    use std::sync::atomic::{AtomicU64, Ordering};
190    pub static ATTN_FWD: AtomicU64 = AtomicU64::new(0);
191    pub static FFN_FWD: AtomicU64 = AtomicU64::new(0);
192    pub static BWD: AtomicU64 = AtomicU64::new(0);
193    pub static GEMM: AtomicU64 = AtomicU64::new(0);
194    pub static GEMM_CALLS: AtomicU64 = AtomicU64::new(0);
195    #[inline]
196    pub fn add(c: &AtomicU64, t: std::time::Instant) {
197        c.fetch_add(t.elapsed().as_nanos() as u64, Ordering::Relaxed);
198    }
199    /// Read and reset, in seconds.
200    pub fn take(c: &AtomicU64) -> f64 {
201        c.swap(0, Ordering::Relaxed) as f64 / 1e9
202    }
203
204    use std::sync::Mutex;
205    /// (n, k, m) → (calls, ns). A handful of shapes dominate a training
206    /// step; naming them decides which ones are worth a GPU submit and
207    /// which are cheaper computed in place.
208    pub static SHAPES: Mutex<Vec<((usize, usize, usize), (u64, u64))>> = Mutex::new(Vec::new());
209    pub fn gemm_shape(n: usize, k: usize, m: usize, t0: std::time::Instant) {
210        let ns = t0.elapsed().as_nanos() as u64;
211        let mut g = SHAPES.lock().unwrap();
212        match g.iter_mut().find(|(s, _)| *s == (n, k, m)) {
213            Some((_, (c, tt))) => {
214                *c += 1;
215                *tt += ns;
216            }
217            None => g.push(((n, k, m), (1, ns))),
218        }
219    }
220    pub fn shape_report(top: usize) -> String {
221        let mut g = SHAPES.lock().unwrap();
222        g.sort_by_key(|(_, (_, ns))| std::cmp::Reverse(*ns));
223        let out = g
224            .iter()
225            .take(top)
226            .map(|((n, k, m), (c, ns))| {
227                format!(
228                    "    [{n}x{k}x{m}] {c} calls, {:.1}s total, {:.1} ms/call",
229                    *ns as f64 / 1e9,
230                    *ns as f64 / 1e6 / *c as f64
231                )
232            })
233            .collect::<Vec<_>>()
234            .join("\n");
235        g.clear();
236        out
237    }
238}
239
240/// Frozen attention operator of one layer — the per-layer dispatch
241/// point for through-backwards (docs/RUST_FCD.md §3).
242enum FcdAttn {
243    Full {
244        wq: Vec<f32>,
245        wk: Vec<f32>,
246        wv: Vec<f32>,
247        /// `[wq; wk; wv]` rows concatenated: one GEMM submit instead of
248        /// three. The profiler priced a submit at ~20 ms on the stand,
249        /// which is 5-10x the multiply it carries at these shapes.
250        wqkv: Vec<f32>,
251        wo: Vec<f32>,
252        q_norm: Option<Vec<f32>>,
253        k_norm: Option<Vec<f32>>,
254        bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
255        /// Qwen3.5: wq rows = 2·nh·hd, per-head [q; gate]; the head
256        /// outputs are multiplied by σ(gate) before o_proj.
257        output_gate: bool,
258    },
259    /// GatedDeltaNet (Qwen3.5 hybrids): never converted, never trained;
260    /// through-backward = BPTT over the window (fcd_ops::gdn_*).
261    Gdn {
262        wqkv: Vec<f32>,
263        wz: Vec<f32>,
264        wa: Vec<f32>,
265        wb: Vec<f32>,
266        conv: Vec<f32>,
267        a_log: Vec<f32>,
268        dt_bias: Vec<f32>,
269        norm: Vec<f32>,
270        wout: Vec<f32>,
271    },
272}
273
274pub(crate) struct FcdLayer {
275    attn: FcdAttn,
276    pub(crate) inter: usize,
277    // Frozen originals: the teacher's LN/FFN and the student's init.
278    pub(crate) iln: Vec<f32>,
279    pub(crate) pln: Vec<f32>,
280    pub(crate) gate: Vec<f32>,
281    pub(crate) up: Vec<f32>,
282    pub(crate) down: Vec<f32>,
283    /// `[gate; up]` prebuilt for the one-submit FFN projection.
284    pub(crate) gu: Vec<f32>,
285}
286
287/// GDN geometry shared by every linear layer (arch.linear_* fields).
288#[derive(Clone, Copy)]
289struct GdnDims {
290    nv: usize,
291    nk: usize,
292    dk: usize,
293    dv: usize,
294    kk: usize,
295}
296
297impl GdnDims {
298    fn c_dim(&self) -> usize {
299        2 * self.nk * self.dk + self.nv * self.dv
300    }
301    fn vd(&self) -> usize {
302        self.nv * self.dv
303    }
304}
305
306/// The f32 training replica of a .cmf model (≤ 1B targets).
307pub struct FcdModel {
308    pub hidden: usize,
309    pub nh: usize,
310    pub nkv: usize,
311    pub hd: usize,
312    pub nl: usize,
313    pub vocab: usize,
314    pub(crate) eps: f64,
315    pub(crate) gemma: bool,
316    rotary_dim: usize,
317    inv_freq: Vec<f64>,
318    /// [vocab, hidden]; also the tied head when `lm_head` is None.
319    pub(crate) embed: Vec<f32>,
320    pub(crate) lm_head: Option<Vec<f32>>,
321    pub(crate) final_norm: Vec<f32>,
322    pub(crate) layers: Vec<FcdLayer>,
323    /// Which layers run the Nyström kernel in the student forward.
324    o1_flags: Vec<bool>,
325    nys: NysCfg,
326    /// GDN geometry (present when the model has linear layers).
327    gdn: Option<GdnDims>,
328    /// Looped Transformer: how many times the layer stack runs per
329    /// token, and whether a final_norm is inserted at each loop
330    /// boundary. A replica that ignores these computes a DIFFERENT
331    /// model — on Nanbeige 4.2 (22 layers x 2 loops) the single-pass
332    /// forward measured perplexity 4773 where the engine reads 16.5,
333    /// and every gradient taken through it was fitted to that fiction.
334    pub(crate) loops: usize,
335    pub(crate) loop_norm: bool,
336    pub(crate) pool: Option<Arc<Pool>>,
337}
338
339fn deq(model: &CmfModel, name: &str) -> Result<Vec<f32>, String> {
340    let e = model
341        .tensor(name)
342        .ok_or_else(|| format!("tensor '{name}' not found"))?;
343    let mut out = vec![0f32; e.n_elems()];
344    cortiq_core::quant::dequant_tensor(e, model.entry_bytes(e), &mut out)?;
345    Ok(out)
346}
347
348impl FcdModel {
349    /// Dequantize a model into the f32 training replica. Refuses what
350    /// the backward cannot honestly differentiate yet (loud, not silent).
351    pub fn from_cmf(model: &CmfModel, o1: &O1Cfg) -> Result<Self, String> {
352        let arch = model.arch().clone();
353        if arch.hidden_act != "silu" {
354            return Err(format!(
355                "fcd/skill-bake: hidden_act '{}' not supported yet (SiLU only)",
356                arch.hidden_act
357            ));
358        }
359        let has_linear = arch
360            .layer_types
361            .iter()
362            .any(|t| matches!(t, LayerType::LinearAttention));
363        let gdn = if has_linear {
364            let lc = arch
365                .linear_core
366                .as_ref()
367                .ok_or_else(|| "model has linear layers but no arch.linear_core".to_string())?;
368            if lc.kind != "gated_delta_net" {
369                return Err(format!(
370                    "linear core '{}' has no FCD backward (only gated_delta_net)",
371                    lc.kind
372                ));
373            }
374            Some(GdnDims {
375                nv: lc.num_heads,
376                nk: arch
377                    .linear_num_key_heads
378                    .ok_or("linear core needs arch.linear_num_key_heads")?,
379                dk: arch
380                    .linear_key_head_dim
381                    .ok_or("linear core needs arch.linear_key_head_dim")?,
382                dv: lc.value_head_dim,
383                kk: arch
384                    .linear_conv_kernel_dim
385                    .ok_or("linear core needs arch.linear_conv_kernel_dim")?,
386            })
387        } else {
388            None
389        };
390        let (nh, nkv, hd, h) = (
391            arch.num_attention_heads,
392            arch.num_kv_heads,
393            arch.head_dim,
394            arch.hidden_size,
395        );
396        let embed = deq(model, "model.embed_tokens.weight")?;
397        let lm_head = if model.tensor("lm_head.weight").is_some() {
398            Some(deq(model, "lm_head.weight")?)
399        } else if arch.tie_word_embeddings {
400            None
401        } else {
402            return Err("no lm_head.weight and tie_word_embeddings is false".into());
403        };
404        let final_norm = deq(model, "model.norm.weight")?;
405
406        let mut layers = Vec::with_capacity(arch.num_layers);
407        for li in 0..arch.num_layers {
408            let p = format!("model.layers.{li}.");
409            if model.tensor(&format!("{p}mlp.gate.weight")).is_some() {
410                return Err(format!(
411                    "layer {li} is MoE — FCD polish supports dense FFN only"
412                ));
413            }
414            let attn = match arch.layer_types.get(li) {
415                Some(LayerType::LinearAttention) => {
416                    let la = |n: &str| deq(model, &format!("{p}linear_attn.{n}"));
417                    FcdAttn::Gdn {
418                        wqkv: la("in_proj_qkv.weight")?,
419                        wz: la("in_proj_z.weight")?,
420                        wa: la("in_proj_a.weight")?,
421                        wb: la("in_proj_b.weight")?,
422                        conv: la("conv1d.weight")?,
423                        a_log: la("A_log")?,
424                        dt_bias: la("dt_bias")?,
425                        norm: la("norm.weight")?,
426                        wout: la("out_proj.weight")?,
427                    }
428                }
429                _ => {
430                    let wq = deq(model, &format!("{p}self_attn.q_proj.weight"))?;
431                    let output_gate = wq.len() == 2 * nh * hd * h;
432                    let opt = |n: &str| -> Option<Vec<f32>> {
433                        model
434                            .tensor(&format!("{p}self_attn.{n}"))
435                            .and_then(|_| deq(model, &format!("{p}self_attn.{n}")).ok())
436                    };
437                    let bias = match (opt("q_proj.bias"), opt("k_proj.bias"), opt("v_proj.bias")) {
438                        (Some(a), Some(b), Some(c)) => Some((a, b, c)),
439                        _ => None,
440                    };
441                    let wk = deq(model, &format!("{p}self_attn.k_proj.weight"))?;
442                    let wv = deq(model, &format!("{p}self_attn.v_proj.weight"))?;
443                    let mut wqkv = Vec::with_capacity(wq.len() + wk.len() + wv.len());
444                    wqkv.extend_from_slice(&wq);
445                    wqkv.extend_from_slice(&wk);
446                    wqkv.extend_from_slice(&wv);
447                    FcdAttn::Full {
448                        wq,
449                        wk,
450                        wv,
451                        wqkv,
452                        wo: deq(model, &format!("{p}self_attn.o_proj.weight"))?,
453                        q_norm: opt("q_norm.weight"),
454                        k_norm: opt("k_norm.weight"),
455                        bias,
456                        output_gate,
457                    }
458                }
459            };
460            let gate = deq(model, &format!("{p}mlp.gate_proj.weight"))?;
461            let inter = gate.len() / h;
462            let up = deq(model, &format!("{p}mlp.up_proj.weight"))?;
463            let mut gu = Vec::with_capacity(gate.len() + up.len());
464            gu.extend_from_slice(&gate);
465            gu.extend_from_slice(&up);
466            layers.push(FcdLayer {
467                attn,
468                inter,
469                iln: deq(model, &format!("{p}input_layernorm.weight"))?,
470                pln: deq(model, &format!("{p}post_attention_layernorm.weight"))?,
471                gate,
472                up,
473                down: deq(model, &format!("{p}mlp.down_proj.weight"))?,
474                gu,
475            });
476        }
477
478        let rotary_dim = ((hd as f32 * arch.partial_rotary_factor) as usize)
479            .max(2)
480            .min(hd);
481        let base = arch.rope_theta;
482        let inv_freq: Vec<f64> = (0..rotary_dim / 2)
483            .map(|i| 1.0 / base.powf(2.0 * i as f64 / rotary_dim as f64))
484            .collect();
485        let loops = arch.num_loops.max(1);
486        let loop_norm = arch.loop_final_norm;
487        let mut flags = o1.layer_flags(arch.num_layers);
488        flags.resize(arch.num_layers, false);
489        // Only full-attention layers are o1-convertible (same rule as
490        // Pipeline::set_o1) — a GDN layer keeps its own operator.
491        for (li, f) in flags.iter_mut().enumerate() {
492            if *f && !matches!(layers[li].attn, FcdAttn::Full { .. }) {
493                *f = false;
494            }
495        }
496        Ok(Self {
497            hidden: h,
498            nh,
499            nkv,
500            hd,
501            nl: arch.num_layers,
502            vocab: arch.vocab_size.min(embed.len() / h),
503            eps: arch.rms_norm_eps,
504            gemma: matches!(arch.norm_style, NormStyle::Gemma),
505            rotary_dim,
506            inv_freq,
507            embed,
508            lm_head,
509            final_norm,
510            layers,
511            o1_flags: flags,
512            // prefill: None = half the window, the same seal point
513            // `cortiq ppl --o1` defaults to (see NysCfg::prefill).
514            nys: NysCfg {
515                m: o1.m,
516                w: o1.w,
517                sink: o1.sink,
518                prefill: None,
519            },
520            gdn,
521            loops,
522            loop_norm,
523            pool: Pool::from_env(),
524        })
525    }
526
527    /// Converted (trainable) layer indices.
528    pub fn converted(&self) -> Vec<usize> {
529        (0..self.nl).filter(|&i| self.o1_flags[i]).collect()
530    }
531
532    fn head_weight(&self) -> &[f32] {
533        self.lm_head.as_deref().unwrap_or(&self.embed)
534    }
535}
536
537// ───────────────────────── trainable state ─────────────────────────
538
539/// Per converted layer, in this fixed order.
540const PARAMS_PER_LAYER: usize = 5; // iln, pln, gate, up, down
541
542/// Master copies + grads + AdamW moments of the trainable tensors.
543pub struct TrainState {
544    pub layers: Vec<usize>,
545    /// layers.len()·5 tensors, layer-major, [iln, pln, gate, up, down].
546    pub data: Vec<Vec<f32>>,
547    grad: Vec<Vec<f32>>,
548    m1: Vec<Vec<f32>>,
549    m2: Vec<Vec<f32>>,
550    step_t: u64,
551}
552
553impl TrainState {
554    pub fn new(fm: &FcdModel) -> Self {
555        let layers = fm.converted();
556        let mut data = Vec::with_capacity(layers.len() * PARAMS_PER_LAYER);
557        for &li in &layers {
558            let l = &fm.layers[li];
559            data.push(l.iln.clone());
560            data.push(l.pln.clone());
561            data.push(l.gate.clone());
562            data.push(l.up.clone());
563            data.push(l.down.clone());
564        }
565        let zeros: Vec<Vec<f32>> = data.iter().map(|d| vec![0f32; d.len()]).collect();
566        Self {
567            layers,
568            grad: zeros.clone(),
569            m1: zeros.clone(),
570            m2: zeros,
571            data,
572            step_t: 0,
573        }
574    }
575
576    fn slot(&self, li: usize) -> Option<usize> {
577        self.layers.iter().position(|&x| x == li)
578    }
579
580    /// Read access to the accumulated gradients (gradcheck harness).
581    #[doc(hidden)]
582    pub fn grads(&self) -> &[Vec<f32>] {
583        &self.grad
584    }
585
586    fn zero_grad(&mut self) {
587        for g in &mut self.grad {
588            for v in g.iter_mut() {
589                *v = 0.0;
590            }
591        }
592    }
593
594    /// Global-norm clip (1.0) + one AdamW step (torch defaults,
595    /// decoupled weight decay).
596    fn clip_and_step(&mut self, lr: f64) -> f64 {
597        let mut sq = 0f64;
598        for g in &self.grad {
599            for &v in g {
600                sq += (v as f64) * (v as f64);
601            }
602        }
603        let gn = sq.sqrt();
604        let scale = if gn > 1.0 { 1.0 / (gn + 1e-6) } else { 1.0 };
605        self.step_t += 1;
606        let bc1 = 1.0 - ADAM_B1.powi(self.step_t as i32);
607        let bc2 = 1.0 - ADAM_B2.powi(self.step_t as i32);
608        for p in 0..self.data.len() {
609            let (d, g, m, v) = (
610                &mut self.data[p],
611                &self.grad[p],
612                &mut self.m1[p],
613                &mut self.m2[p],
614            );
615            for i in 0..d.len() {
616                let gi = g[i] as f64 * scale;
617                let mi = ADAM_B1 * m[i] as f64 + (1.0 - ADAM_B1) * gi;
618                let vi = ADAM_B2 * v[i] as f64 + (1.0 - ADAM_B2) * gi * gi;
619                m[i] = mi as f32;
620                v[i] = vi as f32;
621                let upd = (mi / bc1) / ((vi / bc2).sqrt() + ADAM_EPS) + ADAM_WD * d[i] as f64;
622                d[i] = (d[i] as f64 - lr * upd) as f32;
623            }
624        }
625        gn
626    }
627}
628
629/// LN/FFN weight view of one layer — frozen originals for the teacher
630/// (and non-converted student layers), master copies for trainables.
631#[derive(Clone, Copy)]
632pub(crate) struct LnFfn<'a> {
633    pub(crate) iln: &'a [f32],
634    pub(crate) pln: &'a [f32],
635    pub(crate) gate: &'a [f32],
636    pub(crate) up: &'a [f32],
637    pub(crate) down: &'a [f32],
638    /// `[gate; up]` rows concatenated — one GEMM submit instead of two.
639    /// `None` for Phase-B trained copies, whose rows move every step;
640    /// the frozen majority carries the concat prebuilt at load.
641    pub(crate) gu: Option<&'a [f32]>,
642}
643
644fn ln_ffn<'a>(fm: &'a FcdModel, ts: Option<&'a TrainState>, li: usize) -> LnFfn<'a> {
645    if let Some(t) = ts {
646        if let Some(s) = t.slot(li) {
647            let b = s * PARAMS_PER_LAYER;
648            return LnFfn {
649                iln: &t.data[b],
650                pln: &t.data[b + 1],
651                gate: &t.data[b + 2],
652                up: &t.data[b + 3],
653                down: &t.data[b + 4],
654                gu: None,
655            };
656        }
657    }
658    let l = &fm.layers[li];
659    LnFfn {
660        iln: &l.iln,
661        pln: &l.pln,
662        gate: &l.gate,
663        up: &l.up,
664        down: &l.down,
665        gu: Some(&l.gu),
666    }
667}
668
669// ───────────────────── layer forward (+ recompute) ─────────────────────
670
671/// Intra-layer activations rebuilt during the checkpointed backward.
672enum AttnActs {
673    Full {
674        qpre: Vec<f32>,
675        kpre: Vec<f32>,
676        vproj: Vec<f32>,
677        qrot: Vec<f32>,
678        krot: Vec<f32>,
679        qinv: Vec<f32>,
680        kinv: Vec<f32>,
681        /// Pre-gate per-head attention outputs (needed for the output
682        /// gate's backward); always kept — transient per layer.
683        ao: Vec<f32>,
684        /// Raw gate half of q_proj (empty without an output gate).
685        gate_pre: Vec<f32>,
686    },
687    /// Raw projection streams — the GDN backward replays conv + the
688    /// recurrence from these.
689    Gdn {
690        qkv: Vec<f32>,
691        z: Vec<f32>,
692        a: Vec<f32>,
693        b: Vec<f32>,
694    },
695}
696
697pub(crate) struct LayerActs {
698    inv1: Vec<f32>,
699    attn: AttnActs,
700    pub(crate) h1: Vec<f32>,
701    pub(crate) n2: Vec<f32>,
702    pub(crate) inv2: Vec<f32>,
703    pub(crate) gpre: Vec<f32>,
704    pub(crate) upre: Vec<f32>,
705    pub(crate) act: Vec<f32>,
706}
707
708/// Disjoint-write pointer for pooled per-head scatter (pipeline pattern).
709struct SendMut<T>(*mut T);
710unsafe impl<T> Send for SendMut<T> {}
711unsafe impl<T> Sync for SendMut<T> {}
712impl<T> SendMut<T> {
713    #[inline]
714    unsafe fn at(&self, i: usize) -> *mut T {
715        unsafe { self.0.add(i) }
716    }
717}
718
719impl FcdModel {
720    /// Per-head RMS-norm (qk-norm) + partial RoPE for all rows of a
721    /// projection buffer. `heads` per row, `x` is `[n, heads·hd]`.
722    /// Saves the per-(row, head) rms inv when a norm gain is present.
723    fn qk_norm_rope(
724        &self,
725        x: &mut [f32],
726        norm: Option<&[f32]>,
727        heads: usize,
728        t: usize,
729        inv_out: &mut [f32],
730    ) {
731        let hd = self.hd;
732        let n = x.len() / (heads * hd);
733        for r in 0..n {
734            let pos = r % t;
735            for hh in 0..heads {
736                let s = (r * heads + hh) * hd;
737                let head = &mut x[s..s + hd];
738                if let Some(w) = norm {
739                    let mut inv = [0f32; 1];
740                    let mut y = [0f32; 256];
741                    debug_assert!(hd <= 256);
742                    ops::rmsnorm_fwd(head, w, self.eps, self.gemma, &mut y[..hd], &mut inv);
743                    head.copy_from_slice(&y[..hd]);
744                    inv_out[r * heads + hh] = inv[0];
745                }
746                ops::rope_fwd(&mut head[..self.rotary_dim], pos, &self.inv_freq);
747            }
748        }
749    }
750
751    /// One layer forward over `b` sequences of length `t` (rows are
752    /// b-major). `nystrom` switches converted (Full) student layers to
753    /// the certified matrix kernel (f64 per head); exact heads run in
754    /// f32; GDN layers run the frozen BPTT-capable operator. Returns
755    /// (h_out, intra-layer activations when `want_acts`).
756    #[allow(clippy::too_many_arguments)]
757    fn layer_forward(
758        &self,
759        li: usize,
760        h_in: &[f32],
761        b: usize,
762        t: usize,
763        wts: &LnFfn,
764        nystrom: bool,
765        want_acts: bool,
766    ) -> (Vec<f32>, Option<LayerActs>) {
767        self.layer_forward_scaled(li, h_in, b, t, wts, nystrom, want_acts, None)
768    }
769
770    /// `layer_forward` with an optional per-neuron FFN activation scale
771    /// (the DTG-MA mask σ(m), Patent 2): `act·scale` feeds down_proj.
772    /// `LayerActs.act` stays PRE-scale so the mask backward can read it.
773    #[allow(clippy::too_many_arguments)]
774    pub(crate) fn layer_forward_scaled(
775        &self,
776        li: usize,
777        h_in: &[f32],
778        b: usize,
779        t: usize,
780        wts: &LnFfn,
781        nystrom: bool,
782        want_acts: bool,
783        ffn_scale: Option<&[f32]>,
784    ) -> (Vec<f32>, Option<LayerActs>) {
785        let hsz = self.hidden;
786        let n = b * t;
787        let l = &self.layers[li];
788        let pool = self.pool.as_deref();
789
790        let mut n1 = vec![0f32; n * hsz];
791        let mut inv1 = vec![0f32; n];
792        ops::rmsnorm_fwd(h_in, wts.iln, self.eps, self.gemma, &mut n1, &mut inv1);
793
794        let t_attn = std::time::Instant::now();
795        let (attn_out, attn_acts) = match &l.attn {
796            FcdAttn::Full { .. } => self.full_attn_fwd(&l.attn, &n1, b, t, nystrom),
797            FcdAttn::Gdn { .. } => self.gdn_attn_fwd(&l.attn, &n1, b, t),
798        };
799        prof::add(&prof::ATTN_FWD, t_attn);
800        let t_ffn = std::time::Instant::now();
801
802        let mut h1 = h_in.to_vec();
803        for (a, &x) in h1.iter_mut().zip(&attn_out) {
804            *a += x;
805        }
806
807        let mut n2 = vec![0f32; n * hsz];
808        let mut inv2 = vec![0f32; n];
809        ops::rmsnorm_fwd(&h1, wts.pln, self.eps, self.gemma, &mut n2, &mut inv2);
810
811        let inter = l.inter;
812        let mut gpre = vec![0f32; n * inter];
813        let mut upre = vec![0f32; n * inter];
814        if let Some(gu) = wts.gu {
815            // One submit; split back so everything downstream is
816            // byte-identical to the two-call path.
817            let mut both = vec![0f32; n * 2 * inter];
818            ops::gemm_nt(&n2, gu, &mut both, n, hsz, 2 * inter, pool);
819            for r in 0..n {
820                let row = &both[r * 2 * inter..(r + 1) * 2 * inter];
821                gpre[r * inter..(r + 1) * inter].copy_from_slice(&row[..inter]);
822                upre[r * inter..(r + 1) * inter].copy_from_slice(&row[inter..]);
823            }
824        } else {
825            ops::gemm_nt(&n2, wts.gate, &mut gpre, n, hsz, inter, pool);
826            ops::gemm_nt(&n2, wts.up, &mut upre, n, hsz, inter, pool);
827        }
828        let mut act = vec![0f32; n * inter];
829        for i in 0..n * inter {
830            act[i] = ops::silu(gpre[i]) * upre[i];
831        }
832        let mut ffn = vec![0f32; n * hsz];
833        match ffn_scale {
834            Some(g) => {
835                debug_assert_eq!(g.len(), inter);
836                let mut act2 = act.clone();
837                for r in 0..n {
838                    for (a, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
839                        *a *= gv;
840                    }
841                }
842                ops::gemm_nt(&act2, wts.down, &mut ffn, n, inter, hsz, pool);
843            }
844            None => ops::gemm_nt(&act, wts.down, &mut ffn, n, inter, hsz, pool),
845        }
846        let mut h2 = h1.clone();
847        for (a, &x) in h2.iter_mut().zip(&ffn) {
848            *a += x;
849        }
850
851        let acts = want_acts.then_some(LayerActs {
852            inv1,
853            attn: attn_acts,
854            h1,
855            n2,
856            inv2,
857            gpre,
858            upre,
859            act,
860        });
861        prof::add(&prof::FFN_FWD, t_ffn);
862        (h2, acts)
863    }
864
865    /// Full-attention forward: projections (+optional biases), optional
866    /// per-head [q; gate] split (Qwen3.5 output gate), qk-norm + RoPE,
867    /// per-head exact-or-Nyström attention, σ(gate) multiply, o_proj.
868    fn full_attn_fwd(
869        &self,
870        attn: &FcdAttn,
871        n1: &[f32],
872        b: usize,
873        t: usize,
874        nystrom: bool,
875    ) -> (Vec<f32>, AttnActs) {
876        let FcdAttn::Full {
877            wq,
878            wk,
879            wv,
880            wqkv,
881            wo,
882            q_norm,
883            k_norm,
884            bias,
885            output_gate,
886        } = attn
887        else {
888            unreachable!("full_attn_fwd on a non-Full layer");
889        };
890        let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
891        let n = b * t;
892        let pool = self.pool.as_deref();
893        let qdim = nh * hd;
894        let kvdim = nkv * hd;
895        let rep = nh / nkv;
896        let qrows = if *output_gate { 2 * qdim } else { qdim };
897
898        // One fused submit; the split back into three buffers is a
899        // memcpy that costs microseconds and keeps everything below
900        // this line byte-identical to the unfused path.
901        let fused = qrows + 2 * kvdim;
902        let mut qkv = vec![0f32; n * fused];
903        ops::gemm_nt(n1, wqkv, &mut qkv, n, hsz, fused, pool);
904        let _ = (wq, wk, wv);
905        let mut qraw = vec![0f32; n * qrows];
906        let mut kpre = vec![0f32; n * kvdim];
907        let mut vproj = vec![0f32; n * kvdim];
908        for r in 0..n {
909            let row = &qkv[r * fused..(r + 1) * fused];
910            qraw[r * qrows..(r + 1) * qrows].copy_from_slice(&row[..qrows]);
911            kpre[r * kvdim..(r + 1) * kvdim].copy_from_slice(&row[qrows..qrows + kvdim]);
912            vproj[r * kvdim..(r + 1) * kvdim].copy_from_slice(&row[qrows + kvdim..]);
913        }
914        if let Some((bq, bk, bv)) = bias {
915            for r in 0..n {
916                for (x, bb) in qraw[r * qrows..(r + 1) * qrows].iter_mut().zip(bq) {
917                    *x += bb;
918                }
919                for (x, bb) in kpre[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bk) {
920                    *x += bb;
921                }
922                for (x, bb) in vproj[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bv) {
923                    *x += bb;
924                }
925            }
926        }
927        // Gate split: per-head [q(hd); gate(hd)] (runtime convention).
928        let (qpre, gate_pre) = if *output_gate {
929            let mut qh = vec![0f32; n * qdim];
930            let mut gp = vec![0f32; n * qdim];
931            for r in 0..n {
932                for h in 0..nh {
933                    let src = r * qrows + 2 * h * hd;
934                    let dst = r * qdim + h * hd;
935                    qh[dst..dst + hd].copy_from_slice(&qraw[src..src + hd]);
936                    gp[dst..dst + hd].copy_from_slice(&qraw[src + hd..src + 2 * hd]);
937                }
938            }
939            (qh, gp)
940        } else {
941            (qraw, Vec::new())
942        };
943
944        let mut qrot = qpre.clone();
945        let mut krot = kpre.clone();
946        let mut qinv = vec![0f32; n * nh];
947        let mut kinv = vec![0f32; n * nkv];
948        self.qk_norm_rope(&mut qrot, q_norm.as_deref(), nh, t, &mut qinv);
949        self.qk_norm_rope(&mut krot, k_norm.as_deref(), nkv, t, &mut kinv);
950
951        // ── attention heads: parallel over (sequence, head) ──
952        let mut ao = vec![0f32; n * qdim];
953        {
954            let units = b * nh;
955            let aop = SendMut(ao.as_mut_ptr());
956            let qr = &qrot;
957            let kr = &krot;
958            let vr = &vproj;
959            let nys = self.nys;
960            let run_unit = |u: usize| {
961                let (bi, h) = (u / nh, u % nh);
962                let g = h / rep;
963                if nystrom {
964                    // Certified matrix kernel in f64 (docs/RUST_FCD.md §2.2).
965                    let mut q64 = vec![0f64; t * hd];
966                    let mut k64 = vec![0f64; t * hd];
967                    let mut v64 = vec![0f64; t * hd];
968                    for p in 0..t {
969                        let r = bi * t + p;
970                        for c in 0..hd {
971                            q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
972                            k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
973                            v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
974                        }
975                    }
976                    let mut o64 = vec![0f64; t * hd];
977                    ops::nystrom_head_fwd(&q64, &k64, &v64, t, hd, hd, &nys, &mut o64);
978                    for p in 0..t {
979                        let r = bi * t + p;
980                        for c in 0..hd {
981                            // SAFETY: (row, head) slices are disjoint per unit.
982                            unsafe {
983                                *aop.at(r * qdim + h * hd + c) = o64[p * hd + c] as f32;
984                            }
985                        }
986                    }
987                } else {
988                    let mut q32 = vec![0f32; t * hd];
989                    let mut k32 = vec![0f32; t * hd];
990                    let mut v32 = vec![0f32; t * hd];
991                    for p in 0..t {
992                        let r = bi * t + p;
993                        q32[p * hd..(p + 1) * hd]
994                            .copy_from_slice(&qr[r * qdim + h * hd..r * qdim + (h + 1) * hd]);
995                        k32[p * hd..(p + 1) * hd]
996                            .copy_from_slice(&kr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
997                        v32[p * hd..(p + 1) * hd]
998                            .copy_from_slice(&vr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
999                    }
1000                    let mut o32 = vec![0f32; t * hd];
1001                    ops::attn_head_fwd(&q32, &k32, &v32, t, hd, hd, &mut o32);
1002                    for p in 0..t {
1003                        let r = bi * t + p;
1004                        for c in 0..hd {
1005                            // SAFETY: disjoint (row, head) slices per unit.
1006                            unsafe {
1007                                *aop.at(r * qdim + h * hd + c) = o32[p * hd + c];
1008                            }
1009                        }
1010                    }
1011                }
1012            };
1013            match pool {
1014                Some(p) if units > 1 => p.run(&|widx, nw| {
1015                    for u in (widx..units).step_by(nw) {
1016                        run_unit(u);
1017                    }
1018                }),
1019                _ => {
1020                    for u in 0..units {
1021                        run_unit(u);
1022                    }
1023                }
1024            }
1025        }
1026
1027        // Output gate: multiply the head outputs by σ(gate) before o_proj.
1028        let ao_eff: Vec<f32> = if *output_gate {
1029            ao.iter()
1030                .zip(&gate_pre)
1031                .map(|(&a, &g)| a * (1.0 / (1.0 + (-g).exp())))
1032                .collect()
1033        } else {
1034            ao.clone()
1035        };
1036        let mut attn_out = vec![0f32; n * hsz];
1037        ops::gemm_nt(&ao_eff, wo, &mut attn_out, n, qdim, hsz, pool);
1038        (
1039            attn_out,
1040            AttnActs::Full {
1041                qpre,
1042                kpre,
1043                vproj,
1044                qrot,
1045                krot,
1046                qinv,
1047                kinv,
1048                ao,
1049                gate_pre,
1050            },
1051        )
1052    }
1053
1054    /// GDN forward (frozen operator, teacher AND student): batched
1055    /// projections → f64 conv+SiLU per sequence → pooled per-(seq,
1056    /// k-head) delta-rule recurrence → out_proj. Matches the runtime
1057    /// `gdn_forward` (parity-tested in fcd_gradcheck).
1058    fn gdn_attn_fwd(&self, attn: &FcdAttn, n1: &[f32], b: usize, t: usize) -> (Vec<f32>, AttnActs) {
1059        let FcdAttn::Gdn {
1060            wqkv,
1061            wz,
1062            wa,
1063            wb,
1064            conv,
1065            a_log,
1066            dt_bias,
1067            norm,
1068            wout,
1069        } = attn
1070        else {
1071            unreachable!("gdn_attn_fwd on a non-GDN layer");
1072        };
1073        let d = self.gdn.expect("gdn layer without gdn dims");
1074        let (hsz, n) = (self.hidden, b * t);
1075        let pool = self.pool.as_deref();
1076        let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
1077
1078        let mut qkv = vec![0f32; n * c_dim];
1079        ops::gemm_nt(n1, wqkv, &mut qkv, n, hsz, c_dim, pool);
1080        let mut z = vec![0f32; n * vd];
1081        ops::gemm_nt(n1, wz, &mut z, n, hsz, vd, pool);
1082        let mut a = vec![0f32; n * nv];
1083        ops::gemm_nt(n1, wa, &mut a, n, hsz, nv, pool);
1084        let mut bstr = vec![0f32; n * nv];
1085        ops::gemm_nt(n1, wb, &mut bstr, n, hsz, nv, pool);
1086
1087        let cfg = ops::GdnSeqCfg {
1088            nv: d.nv,
1089            nk: d.nk,
1090            dk: d.dk,
1091            dv: d.dv,
1092            kk: d.kk,
1093            rms_eps: self.eps,
1094            conv,
1095            a_log,
1096            dt_bias,
1097            norm,
1098        };
1099        // f64 streams (runtime-precision recurrence) + per-seq conv.
1100        let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
1101        let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
1102        let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
1103        let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
1104        let mut pre64 = vec![0f64; n * c_dim];
1105        let mut cq64 = vec![0f64; n * c_dim];
1106        for bi in 0..b {
1107            let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1108            ops::gdn_conv_fwd(
1109                &qkv64[r.clone()],
1110                t,
1111                c_dim,
1112                d.kk,
1113                conv,
1114                &mut pre64[r.clone()],
1115                &mut cq64[r],
1116            );
1117        }
1118        let mut of = vec![0f32; n * vd];
1119        {
1120            let units = b * d.nk;
1121            let rep_v = d.nv / d.nk;
1122            let ofp = SendMut(of.as_mut_ptr());
1123            let (cqr, zr, ar, br) = (&cq64, &z64, &a64, &b64);
1124            let cfg_ref = &cfg;
1125            let run_unit = |u: usize| {
1126                let (bi, ko) = (u / d.nk, u % d.nk);
1127                let mut local = vec![0f64; t * vd];
1128                ops::gdn_group_fwd(
1129                    &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
1130                    &zr[bi * t * vd..(bi + 1) * t * vd],
1131                    &ar[bi * t * nv..(bi + 1) * t * nv],
1132                    &br[bi * t * nv..(bi + 1) * t * nv],
1133                    t,
1134                    cfg_ref,
1135                    ko,
1136                    &mut local,
1137                );
1138                for hh in 0..rep_v {
1139                    let h = ko * rep_v + hh;
1140                    for p in 0..t {
1141                        for dj in 0..d.dv {
1142                            // SAFETY: v-head columns are exclusive per unit.
1143                            unsafe {
1144                                *ofp.at((bi * t + p) * vd + h * d.dv + dj) =
1145                                    local[p * vd + h * d.dv + dj] as f32;
1146                            }
1147                        }
1148                    }
1149                }
1150            };
1151            match pool {
1152                Some(p) if units > 1 => p.run(&|widx, nw| {
1153                    for u in (widx..units).step_by(nw) {
1154                        run_unit(u);
1155                    }
1156                }),
1157                _ => {
1158                    for u in 0..units {
1159                        run_unit(u);
1160                    }
1161                }
1162            }
1163        }
1164        let mut attn_out = vec![0f32; n * hsz];
1165        ops::gemm_nt(&of, wout, &mut attn_out, n, vd, hsz, pool);
1166        (attn_out, AttnActs::Gdn { qkv, z, a, b: bstr })
1167    }
1168
1169    /// One layer backward (docs/RUST_FCD.md §2.3 chain), given the
1170    /// recomputed `acts`. Accumulates trainable grads when `grads` is
1171    /// Some; always produces the through-grad dh_in.
1172    #[allow(clippy::too_many_arguments)]
1173    fn layer_backward(
1174        &self,
1175        li: usize,
1176        h_in: &[f32],
1177        b: usize,
1178        t: usize,
1179        wts: &LnFfn,
1180        nystrom: bool,
1181        acts: &LayerActs,
1182        dh2: &[f32],
1183        mut grads: Option<&mut [Vec<f32>]>,
1184    ) -> Vec<f32> {
1185        let hsz = self.hidden;
1186        let n = b * t;
1187        let l = &self.layers[li];
1188        let pool = self.pool.as_deref();
1189        let inter = l.inter;
1190
1191        // ── FFN backward ──
1192        let mut dact = vec![0f32; n * inter];
1193        ops::gemm_dx(dh2, wts.down, &mut dact, n, inter, hsz, pool);
1194        if let Some(g) = grads.as_deref_mut() {
1195            ops::gemm_dw(dh2, &acts.act, &mut g[4], n, inter, hsz, pool);
1196        }
1197        let mut dg = vec![0f32; n * inter];
1198        let mut du = vec![0f32; n * inter];
1199        for i in 0..n * inter {
1200            dg[i] = dact[i] * acts.upre[i] * ops::silu_bwd(acts.gpre[i]);
1201            du[i] = dact[i] * ops::silu(acts.gpre[i]);
1202        }
1203        let mut dn2 = vec![0f32; n * hsz];
1204        ops::gemm_dx(&dg, wts.gate, &mut dn2, n, hsz, inter, pool);
1205        ops::gemm_dx(&du, wts.up, &mut dn2, n, hsz, inter, pool);
1206        if let Some(g) = grads.as_deref_mut() {
1207            ops::gemm_dw(&dg, &acts.n2, &mut g[2], n, hsz, inter, pool);
1208            ops::gemm_dw(&du, &acts.n2, &mut g[3], n, hsz, inter, pool);
1209        }
1210
1211        let mut dh1 = dh2.to_vec();
1212        ops::rmsnorm_bwd(
1213            &acts.h1,
1214            wts.pln,
1215            &acts.inv2,
1216            &dn2,
1217            self.gemma,
1218            &mut dh1,
1219            grads.as_deref_mut().map(|g| &mut g[1][..]),
1220        );
1221
1222        // ── attention backward (dispatch) → dn1 ──
1223        let dn1 = match &l.attn {
1224            FcdAttn::Full { .. } => self.full_attn_bwd(&l.attn, &acts.attn, &dh1, b, t, nystrom),
1225            FcdAttn::Gdn { .. } => self.gdn_attn_bwd(&l.attn, &acts.attn, &dh1, b, t),
1226        };
1227
1228        let mut dh_in = dh1.clone();
1229        ops::rmsnorm_bwd(
1230            h_in,
1231            wts.iln,
1232            &acts.inv1,
1233            &dn1,
1234            self.gemma,
1235            &mut dh_in,
1236            grads.map(|g| &mut g[0][..]),
1237        );
1238        dh_in
1239    }
1240
1241    /// Full-attention through-backward: o_proj → output gate →
1242    /// per-head attention (exact / Nyström-frozen-M) → RoPE → qk-norm →
1243    /// projections. Frozen weights: dX only.
1244    fn full_attn_bwd(
1245        &self,
1246        attn: &FcdAttn,
1247        acts: &AttnActs,
1248        dattn: &[f32],
1249        b: usize,
1250        t: usize,
1251        nystrom: bool,
1252    ) -> Vec<f32> {
1253        let FcdAttn::Full {
1254            wq,
1255            wqkv,
1256            wk,
1257            wv,
1258            wo,
1259            q_norm,
1260            k_norm,
1261            output_gate,
1262            ..
1263        } = attn
1264        else {
1265            unreachable!("full_attn_bwd on a non-Full layer");
1266        };
1267        let AttnActs::Full {
1268            qpre,
1269            kpre,
1270            vproj,
1271            qrot,
1272            krot,
1273            qinv,
1274            kinv,
1275            ao,
1276            gate_pre,
1277        } = acts
1278        else {
1279            unreachable!("acts mismatch");
1280        };
1281        let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
1282        let n = b * t;
1283        let pool = self.pool.as_deref();
1284        let qdim = nh * hd;
1285        let kvdim = nkv * hd;
1286        let rep = nh / nkv;
1287        let qrows = if *output_gate { 2 * qdim } else { qdim };
1288
1289        let mut dao_eff = vec![0f32; n * qdim];
1290        ops::gemm_dx(dattn, wo, &mut dao_eff, n, qdim, hsz, pool);
1291        // Output gate: ao_eff = ao·σ(g) → dao = d·σ(g), dg = d·ao·σ′(g).
1292        let (dao, dgate) = if *output_gate {
1293            let mut dao = vec![0f32; n * qdim];
1294            let mut dgp = vec![0f32; n * qdim];
1295            for i in 0..n * qdim {
1296                let sig = 1.0 / (1.0 + (-gate_pre[i]).exp());
1297                dao[i] = dao_eff[i] * sig;
1298                dgp[i] = dao_eff[i] * ao[i] * sig * (1.0 - sig);
1299            }
1300            (dao, dgp)
1301        } else {
1302            (dao_eff, Vec::new())
1303        };
1304
1305        let mut dqrot = vec![0f32; n * qdim];
1306        let mut dkrot = vec![0f32; n * kvdim];
1307        let mut dvproj = vec![0f32; n * kvdim];
1308        {
1309            // Parallel over (sequence, kv-group): a unit owns the dk/dv
1310            // slices of its group and the dq slices of its rep Q heads.
1311            let units = b * nkv;
1312            let dqp = SendMut(dqrot.as_mut_ptr());
1313            let dkp = SendMut(dkrot.as_mut_ptr());
1314            let dvp = SendMut(dvproj.as_mut_ptr());
1315            let (qr, kr, vr) = (qrot, krot, vproj);
1316            let daor = &dao;
1317            let nys = self.nys;
1318            let run_unit = |u: usize| {
1319                let (bi, g) = (u / nkv, u % nkv);
1320                let mut k64 = vec![0f64; t * hd];
1321                let mut v64 = vec![0f64; t * hd];
1322                for p in 0..t {
1323                    let r = bi * t + p;
1324                    for c in 0..hd {
1325                        k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
1326                        v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
1327                    }
1328                }
1329                let mut dk64 = vec![0f64; t * hd];
1330                let mut dv64 = vec![0f64; t * hd];
1331                let mut q64 = vec![0f64; t * hd];
1332                let mut do64 = vec![0f64; t * hd];
1333                let mut dq64 = vec![0f64; t * hd];
1334                for hh in 0..rep {
1335                    let h = g * rep + hh;
1336                    for p in 0..t {
1337                        let r = bi * t + p;
1338                        for c in 0..hd {
1339                            q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
1340                            do64[p * hd + c] = daor[r * qdim + h * hd + c] as f64;
1341                        }
1342                    }
1343                    for v in dq64.iter_mut() {
1344                        *v = 0.0;
1345                    }
1346                    if nystrom {
1347                        ops::nystrom_head_bwd(
1348                            &q64, &k64, &v64, &do64, t, hd, hd, &nys, &mut dq64, &mut dk64,
1349                            &mut dv64,
1350                        );
1351                    } else {
1352                        ops::attn_head_bwd(
1353                            &q64, &k64, &v64, &do64, t, hd, hd, &mut dq64, &mut dk64, &mut dv64,
1354                        );
1355                    }
1356                    for p in 0..t {
1357                        let r = bi * t + p;
1358                        for c in 0..hd {
1359                            // SAFETY: disjoint (row, head) slices per unit.
1360                            unsafe {
1361                                *dqp.at(r * qdim + h * hd + c) = dq64[p * hd + c] as f32;
1362                            }
1363                        }
1364                    }
1365                }
1366                for p in 0..t {
1367                    let r = bi * t + p;
1368                    for c in 0..hd {
1369                        // SAFETY: disjoint (row, group) slices per unit.
1370                        unsafe {
1371                            *dkp.at(r * kvdim + g * hd + c) = dk64[p * hd + c] as f32;
1372                            *dvp.at(r * kvdim + g * hd + c) = dv64[p * hd + c] as f32;
1373                        }
1374                    }
1375                }
1376            };
1377            match pool {
1378                Some(p) if units > 1 => p.run(&|widx, nw| {
1379                    for u in (widx..units).step_by(nw) {
1380                        run_unit(u);
1381                    }
1382                }),
1383                _ => {
1384                    for u in 0..units {
1385                        run_unit(u);
1386                    }
1387                }
1388            }
1389        }
1390
1391        // qk-norm + RoPE through-grads (frozen gains → no dw).
1392        let mut dqpre = vec![0f32; n * qdim];
1393        let mut dkpre = vec![0f32; n * kvdim];
1394        for r in 0..n {
1395            let pos = r % t;
1396            for h in 0..nh {
1397                let s = r * qdim + h * hd;
1398                ops::rope_bwd(&mut dqrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1399                match q_norm {
1400                    Some(w) => ops::rmsnorm_bwd(
1401                        &qpre[s..s + hd],
1402                        w,
1403                        &qinv[r * nh + h..r * nh + h + 1],
1404                        &dqrot[s..s + hd],
1405                        self.gemma,
1406                        &mut dqpre[s..s + hd],
1407                        None,
1408                    ),
1409                    None => dqpre[s..s + hd].copy_from_slice(&dqrot[s..s + hd]),
1410                }
1411            }
1412            for g in 0..nkv {
1413                let s = r * kvdim + g * hd;
1414                ops::rope_bwd(&mut dkrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1415                match k_norm {
1416                    Some(w) => ops::rmsnorm_bwd(
1417                        &kpre[s..s + hd],
1418                        w,
1419                        &kinv[r * nkv + g..r * nkv + g + 1],
1420                        &dkrot[s..s + hd],
1421                        self.gemma,
1422                        &mut dkpre[s..s + hd],
1423                        None,
1424                    ),
1425                    None => dkpre[s..s + hd].copy_from_slice(&dkrot[s..s + hd]),
1426                }
1427            }
1428        }
1429
1430        // Re-interleave [dq; dgate] per head for gated projections.
1431        let dqraw: Vec<f32> = if *output_gate {
1432            let mut dq = vec![0f32; n * qrows];
1433            for r in 0..n {
1434                for h in 0..nh {
1435                    let dst = r * qrows + 2 * h * hd;
1436                    let src = r * qdim + h * hd;
1437                    dq[dst..dst + hd].copy_from_slice(&dqpre[src..src + hd]);
1438                    dq[dst + hd..dst + 2 * hd].copy_from_slice(&dgate[src..src + hd]);
1439                }
1440            }
1441            dq
1442        } else {
1443            dqpre
1444        };
1445
1446        // Projections (frozen weights → dX only; bias add is identity).
1447        // One fused submit: wqkv's row layout is exactly what gemm_dx
1448        // wants, so the sum dq·Wq + dk·Wk + dv·Wv is a concat of the
1449        // gradients and a single call.
1450        let fused = qrows + 2 * kvdim;
1451        let mut dqkv = vec![0f32; n * fused];
1452        for r in 0..n {
1453            let row = &mut dqkv[r * fused..(r + 1) * fused];
1454            row[..qrows].copy_from_slice(&dqraw[r * qrows..(r + 1) * qrows]);
1455            row[qrows..qrows + kvdim].copy_from_slice(&dkpre[r * kvdim..(r + 1) * kvdim]);
1456            row[qrows + kvdim..].copy_from_slice(&dvproj[r * kvdim..(r + 1) * kvdim]);
1457        }
1458        let mut dn1 = vec![0f32; n * hsz];
1459        ops::gemm_dx(&dqkv, wqkv, &mut dn1, n, hsz, fused, pool);
1460        let _ = (wq, wk, wv);
1461        dn1
1462    }
1463
1464    /// GDN through-backward: out_proj → pooled per-(seq, k-head) BPTT
1465    /// (fcd_ops::gdn_group_bwd) → conv backward → projections. Frozen
1466    /// weights: dX only.
1467    fn gdn_attn_bwd(
1468        &self,
1469        attn: &FcdAttn,
1470        acts: &AttnActs,
1471        dattn: &[f32],
1472        b: usize,
1473        t: usize,
1474    ) -> Vec<f32> {
1475        let FcdAttn::Gdn {
1476            wqkv,
1477            wz,
1478            wa,
1479            wb,
1480            conv,
1481            a_log,
1482            dt_bias,
1483            norm,
1484            wout,
1485        } = attn
1486        else {
1487            unreachable!("gdn_attn_bwd on a non-GDN layer");
1488        };
1489        let AttnActs::Gdn { qkv, z, a, b: bstr } = acts else {
1490            unreachable!("acts mismatch");
1491        };
1492        let d = self.gdn.expect("gdn layer without gdn dims");
1493        let (hsz, n) = (self.hidden, b * t);
1494        let pool = self.pool.as_deref();
1495        let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
1496
1497        let mut dof = vec![0f32; n * vd];
1498        ops::gemm_dx(dattn, wout, &mut dof, n, vd, hsz, pool);
1499
1500        let cfg = ops::GdnSeqCfg {
1501            nv: d.nv,
1502            nk: d.nk,
1503            dk: d.dk,
1504            dv: d.dv,
1505            kk: d.kk,
1506            rms_eps: self.eps,
1507            conv,
1508            a_log,
1509            dt_bias,
1510            norm,
1511        };
1512        let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
1513        let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
1514        let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
1515        let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
1516        let dof64: Vec<f64> = dof.iter().map(|&v| v as f64).collect();
1517        let mut pre64 = vec![0f64; n * c_dim];
1518        let mut cq64 = vec![0f64; n * c_dim];
1519        for bi in 0..b {
1520            let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1521            ops::gdn_conv_fwd(
1522                &qkv64[r.clone()],
1523                t,
1524                c_dim,
1525                d.kk,
1526                conv,
1527                &mut pre64[r.clone()],
1528                &mut cq64[r],
1529            );
1530        }
1531
1532        let mut dcq64 = vec![0f64; n * c_dim];
1533        let mut dz64 = vec![0f64; n * vd];
1534        let mut da64 = vec![0f64; n * nv];
1535        let mut db64 = vec![0f64; n * nv];
1536        {
1537            let units = b * d.nk;
1538            let rep_v = d.nv / d.nk;
1539            let kd = d.nk * d.dk;
1540            let dcqp = SendMut(dcq64.as_mut_ptr());
1541            let dzp = SendMut(dz64.as_mut_ptr());
1542            let dap = SendMut(da64.as_mut_ptr());
1543            let dbp = SendMut(db64.as_mut_ptr());
1544            let (cqr, zr, ar, br, dor) = (&cq64, &z64, &a64, &b64, &dof64);
1545            let cfg_ref = &cfg;
1546            let run_unit = |u: usize| {
1547                let (bi, ko) = (u / d.nk, u % d.nk);
1548                // Full-width locals — the group only fills its own
1549                // channels; the scatter below copies exactly those.
1550                let mut dcq_l = vec![0f64; t * c_dim];
1551                let mut dz_l = vec![0f64; t * vd];
1552                let mut da_l = vec![0f64; t * nv];
1553                let mut db_l = vec![0f64; t * nv];
1554                ops::gdn_group_bwd(
1555                    &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
1556                    &zr[bi * t * vd..(bi + 1) * t * vd],
1557                    &ar[bi * t * nv..(bi + 1) * t * nv],
1558                    &br[bi * t * nv..(bi + 1) * t * nv],
1559                    t,
1560                    cfg_ref,
1561                    ko,
1562                    &dor[bi * t * vd..(bi + 1) * t * vd],
1563                    &mut dcq_l,
1564                    &mut dz_l,
1565                    &mut da_l,
1566                    &mut db_l,
1567                );
1568                // SAFETY of every store below: the written channel /
1569                // column ranges are exclusively owned by (bi, ko).
1570                for p in 0..t {
1571                    let row = (bi * t + p) * c_dim;
1572                    for c in ko * d.dk..(ko + 1) * d.dk {
1573                        unsafe {
1574                            *dcqp.at(row + c) = dcq_l[p * c_dim + c];
1575                            *dcqp.at(row + kd + c) = dcq_l[p * c_dim + kd + c];
1576                        }
1577                    }
1578                    for hh in 0..rep_v {
1579                        let h = ko * rep_v + hh;
1580                        for dj in 0..d.dv {
1581                            unsafe {
1582                                *dcqp.at(row + 2 * kd + h * d.dv + dj) =
1583                                    dcq_l[p * c_dim + 2 * kd + h * d.dv + dj];
1584                                *dzp.at((bi * t + p) * vd + h * d.dv + dj) =
1585                                    dz_l[p * vd + h * d.dv + dj];
1586                            }
1587                        }
1588                        unsafe {
1589                            *dap.at((bi * t + p) * nv + h) = da_l[p * nv + h];
1590                            *dbp.at((bi * t + p) * nv + h) = db_l[p * nv + h];
1591                        }
1592                    }
1593                }
1594            };
1595            match pool {
1596                Some(p) if units > 1 => p.run(&|widx, nw| {
1597                    for u in (widx..units).step_by(nw) {
1598                        run_unit(u);
1599                    }
1600                }),
1601                _ => {
1602                    for u in 0..units {
1603                        run_unit(u);
1604                    }
1605                }
1606            }
1607        }
1608
1609        let mut dqkv64 = vec![0f64; n * c_dim];
1610        for bi in 0..b {
1611            let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1612            ops::gdn_conv_bwd(
1613                &pre64[r.clone()],
1614                t,
1615                c_dim,
1616                d.kk,
1617                conv,
1618                &dcq64[r.clone()],
1619                &mut dqkv64[r],
1620            );
1621        }
1622        let to32 = |v: &[f64]| -> Vec<f32> { v.iter().map(|&x| x as f32).collect() };
1623        let (dqkv, dz, da, db) = (to32(&dqkv64), to32(&dz64), to32(&da64), to32(&db64));
1624
1625        let mut dn1 = vec![0f32; n * hsz];
1626        ops::gemm_dx(&dqkv, wqkv, &mut dn1, n, hsz, c_dim, pool);
1627        ops::gemm_dx(&dz, wz, &mut dn1, n, hsz, vd, pool);
1628        ops::gemm_dx(&da, wa, &mut dn1, n, hsz, nv, pool);
1629        ops::gemm_dx(&db, wb, &mut dn1, n, hsz, nv, pool);
1630        dn1
1631    }
1632
1633    /// Full forward: embeddings → layers → final hidden [b·t, hidden].
1634    /// `student` switches converted layers to the Nyström kernel and
1635    /// reads trainable weights from `ts`; `keep` collects each layer's
1636    /// input hidden for the checkpointed backward.
1637    fn forward_hidden(
1638        &self,
1639        ids: &[u32],
1640        b: usize,
1641        t: usize,
1642        ts: Option<&TrainState>,
1643        student: bool,
1644        mut keep: Option<&mut Vec<Vec<f32>>>,
1645    ) -> Vec<f32> {
1646        let hsz = self.hidden;
1647        let mut h = vec![0f32; b * t * hsz];
1648        for (r, &id) in ids.iter().enumerate() {
1649            let src = (id as usize).min(self.embed.len() / hsz - 1) * hsz;
1650            h[r * hsz..(r + 1) * hsz].copy_from_slice(&self.embed[src..src + hsz]);
1651        }
1652        for li in 0..self.nl {
1653            if let Some(k) = keep.as_deref_mut() {
1654                k.push(h.clone());
1655            }
1656            let wts = ln_ffn(self, if student { ts } else { None }, li);
1657            let nys = student && self.o1_flags[li];
1658            h = self.layer_forward(li, &h, b, t, &wts, nys, false).0;
1659        }
1660        h
1661    }
1662
1663    /// Loss head: chunked tied-lm_head CE+KL against the teacher hidden,
1664    /// returning (ce_mean, kl_mean, dHidden_student).
1665    fn loss_and_dhidden(
1666        &self,
1667        hs: &[f32],
1668        ht: &[f32],
1669        targets: &[u32],
1670        kl_w: f64,
1671    ) -> (f64, f64, Vec<f32>) {
1672        let hsz = self.hidden;
1673        let n = targets.len();
1674        let pool = self.pool.as_deref();
1675        let wh = self.head_weight();
1676        let vs = self.vocab;
1677
1678        let mut ns = vec![0f32; n * hsz];
1679        let mut invs = vec![0f32; n];
1680        ops::rmsnorm_fwd(
1681            hs,
1682            &self.final_norm,
1683            self.eps,
1684            self.gemma,
1685            &mut ns,
1686            &mut invs,
1687        );
1688        let mut nt = vec![0f32; n * hsz];
1689        let mut invt = vec![0f32; n];
1690        ops::rmsnorm_fwd(
1691            ht,
1692            &self.final_norm,
1693            self.eps,
1694            self.gemma,
1695            &mut nt,
1696            &mut invt,
1697        );
1698
1699        let inv_n = 1.0 / n as f64;
1700        let mut ce_sum = 0f64;
1701        let mut kl_sum = 0f64;
1702        let mut dns = vec![0f32; n * hsz];
1703        let mut ls = vec![0f32; LM_CHUNK * vs];
1704        let mut lt = vec![0f32; LM_CHUNK * vs];
1705        let mut dlg = vec![0f32; LM_CHUNK * vs];
1706        let mut r0 = 0usize;
1707        while r0 < n {
1708            let r1 = (r0 + LM_CHUNK).min(n);
1709            let c = r1 - r0;
1710            ops::gemm_nt(
1711                &ns[r0 * hsz..r1 * hsz],
1712                wh,
1713                &mut ls[..c * vs],
1714                c,
1715                hsz,
1716                vs,
1717                pool,
1718            );
1719            ops::gemm_nt(
1720                &nt[r0 * hsz..r1 * hsz],
1721                wh,
1722                &mut lt[..c * vs],
1723                c,
1724                hsz,
1725                vs,
1726                pool,
1727            );
1728            for r in 0..c {
1729                let (ce, kl) = ops::ce_kl_position(
1730                    &ls[r * vs..(r + 1) * vs],
1731                    &lt[r * vs..(r + 1) * vs],
1732                    targets[r0 + r] as usize,
1733                    kl_w,
1734                    inv_n,
1735                    &mut dlg[r * vs..(r + 1) * vs],
1736                );
1737                ce_sum += ce;
1738                kl_sum += kl;
1739            }
1740            ops::gemm_dx(
1741                &dlg[..c * vs],
1742                wh,
1743                &mut dns[r0 * hsz..r1 * hsz],
1744                c,
1745                hsz,
1746                vs,
1747                pool,
1748            );
1749            r0 = r1;
1750        }
1751
1752        let mut dhs = vec![0f32; n * hsz];
1753        ops::rmsnorm_bwd(
1754            hs,
1755            &self.final_norm,
1756            &invs,
1757            &dns,
1758            self.gemma,
1759            &mut dhs,
1760            None,
1761        );
1762        (ce_sum * inv_n, kl_sum * inv_n, dhs)
1763    }
1764
1765    /// Checkpointed backward: per layer, recompute the intra-layer
1766    /// activations and differentiate.
1767    fn backward(
1768        &self,
1769        b: usize,
1770        t: usize,
1771        keep: &[Vec<f32>],
1772        dh_last: Vec<f32>,
1773        ts: &mut TrainState,
1774    ) {
1775        // Split-borrow: the weight view reads `data`, the grads write
1776        // `grad` — disjoint fields of TrainState.
1777        let TrainState {
1778            layers, data, grad, ..
1779        } = ts;
1780        let mut dh = dh_last;
1781        for li in (0..self.nl).rev() {
1782            let h_in = &keep[li];
1783            let nys = self.o1_flags[li];
1784            let slot = layers.iter().position(|&x| x == li);
1785            let wts = match slot {
1786                Some(s) => {
1787                    let bi = s * PARAMS_PER_LAYER;
1788                    LnFfn {
1789                        iln: &data[bi],
1790                        pln: &data[bi + 1],
1791                        gate: &data[bi + 2],
1792                        up: &data[bi + 3],
1793                        down: &data[bi + 4],
1794                        gu: None,
1795                    }
1796                }
1797                None => {
1798                    let l = &self.layers[li];
1799                    LnFfn {
1800                        iln: &l.iln,
1801                        pln: &l.pln,
1802                        gate: &l.gate,
1803                        up: &l.up,
1804                        down: &l.down,
1805                        gu: Some(&l.gu),
1806                    }
1807                }
1808            };
1809            let (_, acts) = self.layer_forward(li, h_in, b, t, &wts, nys, true);
1810            let acts = acts.expect("want_acts");
1811            dh = match slot {
1812                Some(s) => {
1813                    let gb = s * PARAMS_PER_LAYER;
1814                    let gr = &mut grad[gb..gb + PARAMS_PER_LAYER];
1815                    self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, Some(gr))
1816                }
1817                None => self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, None),
1818            };
1819        }
1820    }
1821
1822    /// Test-only: one full training-graph evaluation — teacher forward,
1823    /// student forward, CE+KL loss, checkpointed backward into the
1824    /// grads. Returns the weighted total loss. The block-level
1825    /// gradcheck runs finite differences over trainable weights through
1826    /// this, which exercises EVERY through-grad in the graph (layer-0
1827    /// gains flow through all attention/rope/qk-norm/GQA paths above).
1828    #[doc(hidden)]
1829    pub fn loss_and_grads_for_test(
1830        &self,
1831        ids: &[u32],
1832        tgt: &[u32],
1833        b: usize,
1834        t: usize,
1835        ts: &mut TrainState,
1836        kl_w: f64,
1837    ) -> f64 {
1838        let ht = self.forward_hidden(ids, b, t, None, false, None);
1839        let mut keep = Vec::with_capacity(self.nl);
1840        let hs = self.forward_hidden(ids, b, t, Some(ts), true, Some(&mut keep));
1841        let (ce, kl, dhs) = self.loss_and_dhidden(&hs, &ht, tgt, kl_w);
1842        ts.zero_grad();
1843        self.backward(b, t, &keep, dhs, ts);
1844        (1.0 - kl_w) * ce + kl_w * kl
1845    }
1846
1847    /// Teacher-forced CE perplexity on deterministic evenly-spaced val
1848    /// windows (`heal_hybridk_06b.py::val_ppl` discipline — random
1849    /// windows made gate comparisons ride ±15% noise).
1850    pub fn val_ppl(
1851        &self,
1852        va: &[u32],
1853        ts: Option<&TrainState>,
1854        student: bool,
1855        bs: usize,
1856        nrounds: usize,
1857        seq: usize,
1858    ) -> f64 {
1859        let nwin = nrounds * bs;
1860        if va.len() < seq + 2 || nwin == 0 {
1861            return f64::NAN;
1862        }
1863        let stride = (va.len() - seq - 1) / nwin;
1864        let hsz = self.hidden;
1865        let wh = self.head_weight();
1866        let vs = self.vocab;
1867        let pool = self.pool.as_deref();
1868        let mut nll = 0f64;
1869        let mut cnt = 0usize;
1870        for j in 0..nrounds {
1871            let mut ids = Vec::with_capacity(bs * seq);
1872            let mut tgt = Vec::with_capacity(bs * seq);
1873            for bi in 0..bs {
1874                let off = ((j * bs + bi) * stride.max(1)).min(va.len() - seq - 1);
1875                ids.extend_from_slice(&va[off..off + seq]);
1876                tgt.extend_from_slice(&va[off + 1..off + seq + 1]);
1877            }
1878            let h = self.forward_hidden(&ids, bs, seq, ts, student, None);
1879            let n = bs * seq;
1880            let mut ns = vec![0f32; n * hsz];
1881            let mut inv = vec![0f32; n];
1882            ops::rmsnorm_fwd(
1883                &h,
1884                &self.final_norm,
1885                self.eps,
1886                self.gemma,
1887                &mut ns,
1888                &mut inv,
1889            );
1890            let mut lg = vec![0f32; LM_CHUNK * vs];
1891            let mut r0 = 0usize;
1892            while r0 < n {
1893                let r1 = (r0 + LM_CHUNK).min(n);
1894                let c = r1 - r0;
1895                ops::gemm_nt(
1896                    &ns[r0 * hsz..r1 * hsz],
1897                    wh,
1898                    &mut lg[..c * vs],
1899                    c,
1900                    hsz,
1901                    vs,
1902                    pool,
1903                );
1904                for r in 0..c {
1905                    let row = &lg[r * vs..(r + 1) * vs];
1906                    let target = tgt[r0 + r] as usize;
1907                    let mut mx = f64::NEG_INFINITY;
1908                    for &v in row {
1909                        mx = mx.max(v as f64);
1910                    }
1911                    let mut s = 0f64;
1912                    for &v in row {
1913                        s += (v as f64 - mx).exp();
1914                    }
1915                    nll += mx + s.ln() - row[target.min(vs - 1)] as f64;
1916                    cnt += 1;
1917                }
1918                r0 = r1;
1919            }
1920        }
1921        (nll / cnt.max(1) as f64).exp()
1922    }
1923}
1924
1925// ─────────────────────────── training loop ───────────────────────────
1926
1927/// Run the full certified polish: train, early-stop/restore-best, and
1928/// write `<out>` (source tensors byte-copied, polished LN/FFN as f32).
1929///
1930/// With `gate` (Patent 16 draft, claim 13), every eval checkpoint is
1931/// additionally scored by greedy generation through the REAL streaming
1932/// O(1) runtime, and the restored checkpoint is the lowest-ppl one
1933/// AMONG GATE-PASSERS; if none passes, the zero-shot state is restored
1934/// (identity polish) — the stage never makes generation worse than
1935/// conversion alone.
1936pub fn run_polish(
1937    model: &Arc<CmfModel>,
1938    o1: &O1Cfg,
1939    hp: &FcdHyper,
1940    tr: &[u32],
1941    va: &[u32],
1942    out: &std::path::Path,
1943    gate: Option<&GenGateCfg>,
1944) -> Result<FcdReport, String> {
1945    if tr.len() < hp.seq + 2 {
1946        return Err(format!(
1947            "train corpus too small: {} tokens < seq+2 = {}",
1948            tr.len(),
1949            hp.seq + 2
1950        ));
1951    }
1952    let fm = FcdModel::from_cmf(model, o1)?;
1953    let converted = fm.converted();
1954    if converted.is_empty() {
1955        return Err("no converted layers under this --o1 spec (nothing to polish)".into());
1956    }
1957    tracing::info!(
1958        "fcd: {} layers converted ({} trainable tensors), m={} w={} sink={}, \
1959         corpus train {} / val {} tokens",
1960        converted.len(),
1961        converted.len() * PARAMS_PER_LAYER,
1962        fm.nys.m,
1963        fm.nys.w,
1964        fm.nys.sink,
1965        tr.len(),
1966        va.len()
1967    );
1968
1969    let mut ts = TrainState::new(&fm);
1970    let teacher_ppl = fm.val_ppl(va, None, false, hp.bs, 2, hp.seq);
1971    let ppl_start = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
1972    tracing::info!(
1973        "fcd: quick-val teacher ppl {teacher_ppl:.2} | zero-shot o1 student ppl {ppl_start:.2}"
1974    );
1975
1976    // ── generation gate (claim 13): baseline at step 0 ──
1977    let mut gate_state: Option<(Pipeline, Vec<f64>)> = match gate {
1978        Some(g) if !g.prompts.is_empty() => {
1979            let greedy = SamplerConfig {
1980                temperature: 0.0,
1981                top_p: 1.0,
1982                top_k: 0,
1983                repetition_penalty: 1.0,
1984                min_p: 0.0,
1985                seed: Some(0),
1986                suppress_tokens: Vec::new(),
1987            };
1988            let mut pipe = Pipeline::from_model(model, greedy)
1989                .map_err(|e| format!("gen-gate pipeline: {e}"))?;
1990            pipe.set_o1(Some(o1.clone()));
1991            apply_trainables(&mut pipe, &fm, &ts);
1992            let base = gate_gen_scores(&mut pipe, g)?;
1993            tracing::info!("fcd gen-gate baseline loop-scores: {base:?}");
1994            Some((pipe, base))
1995        }
1996        Some(_) => {
1997            tracing::warn!("fcd gen-gate requested but val stream too short — gate off");
1998            None
1999        }
2000        None => None,
2001    };
2002    // Identity fallback: the pre-training master copies.
2003    let init_snapshot: Option<Vec<Vec<f32>>> = gate_state.is_some().then(|| ts.data.clone());
2004    let mut gate_evals: Vec<(usize, f64, Vec<f64>, bool)> = Vec::new();
2005
2006    let mut rng = SplitMix64::new(hp.seed);
2007    let mut best: (f64, Option<Vec<Vec<f32>>>, usize) = (f64::INFINITY, None, 0);
2008    let mut losses: Vec<(f64, f64)> = Vec::with_capacity(hp.steps);
2009    let t0 = std::time::Instant::now();
2010    let n_per_step = hp.bs * hp.seq;
2011    for st in 1..=hp.steps {
2012        // Fresh random windows each step (the recipe; indices need not
2013        // match the torch RNG — the distribution does).
2014        let mut ids = Vec::with_capacity(n_per_step);
2015        let mut tgt = Vec::with_capacity(n_per_step);
2016        for _ in 0..hp.bs {
2017            let off = (rng.next_u64() as usize) % (tr.len() - hp.seq - 1);
2018            ids.extend_from_slice(&tr[off..off + hp.seq]);
2019            tgt.extend_from_slice(&tr[off + 1..off + hp.seq + 1]);
2020        }
2021
2022        let ht = fm.forward_hidden(&ids, hp.bs, hp.seq, None, false, None);
2023        let mut keep: Vec<Vec<f32>> = Vec::with_capacity(fm.nl);
2024        let hs = fm.forward_hidden(&ids, hp.bs, hp.seq, Some(&ts), true, Some(&mut keep));
2025        let (ce, kl, dhs) = fm.loss_and_dhidden(&hs, &ht, &tgt, hp.kl_w);
2026        ts.zero_grad();
2027        fm.backward(hp.bs, hp.seq, &keep, dhs, &mut ts);
2028        let gn = ts.clip_and_step(hp.lr);
2029        losses.push((ce, kl));
2030
2031        let el = t0.elapsed().as_secs_f64();
2032        tracing::info!(
2033            "fcd step {st}/{}: ce {ce:.3} kl {kl:.3} |g| {gn:.3} ({:.1}s/step)",
2034            hp.steps,
2035            el / st as f64
2036        );
2037        if hp.eval_every > 0 && st % hp.eval_every == 0 {
2038            let p = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
2039            match (&mut gate_state, gate) {
2040                (Some((pipe, base)), Some(g)) => {
2041                    apply_trainables(pipe, &fm, &ts);
2042                    let scores = gate_gen_scores(pipe, g)?;
2043                    let pass = gate_pass(&scores, base, g.threshold, g.baseline_slack);
2044                    let tag = if pass && p < best.0 {
2045                        best = (p, Some(ts.data.clone()), st);
2046                        " *best*"
2047                    } else {
2048                        ""
2049                    };
2050                    tracing::info!(
2051                        "fcd eval step {st}: val ppl {p:.2} | gen-gate {}                          (loop-scores {scores:?}){tag}",
2052                        if pass { "PASS" } else { "FAIL" }
2053                    );
2054                    gate_evals.push((st, p, scores, pass));
2055                }
2056                _ => {
2057                    let tag = if p < best.0 {
2058                        best = (p, Some(ts.data.clone()), st);
2059                        " *best*"
2060                    } else {
2061                        ""
2062                    };
2063                    tracing::info!("fcd eval step {st}: val ppl {p:.2}{tag}");
2064                }
2065            }
2066        }
2067    }
2068
2069    // Early stop: restore the best checkpoint (certified: best was step
2070    // 150 of 300 in the torch run). Under the gate, `best` only ever
2071    // held GATE-PASSING checkpoints; none passing → identity restore.
2072    let mut gate_chosen: Option<usize> = None;
2073    if let Some(snap) = best.1.take() {
2074        ts.data = snap;
2075        gate_chosen = Some(best.2);
2076        tracing::info!(
2077            "fcd: restored best checkpoint from step {} (val ppl {:.2})",
2078            best.2,
2079            best.0
2080        );
2081    } else if let Some(init) = init_snapshot {
2082        ts.data = init;
2083        tracing::info!(
2084            "fcd: polish rejected by generation gate — identity artifact              (zero-shot state written; claim 13 floor)"
2085        );
2086    }
2087    let ppl_final = fm.val_ppl(va, Some(&ts), true, hp.bs, 6, hp.seq);
2088    let report = FcdReport {
2089        converted: converted.clone(),
2090        teacher_ppl,
2091        ppl_start,
2092        ppl_best: best.0.min(ppl_final),
2093        best_step: best.2,
2094        ppl_final,
2095        steps_run: hp.steps,
2096        sec_per_step: t0.elapsed().as_secs_f64() / hp.steps.max(1) as f64,
2097        losses,
2098        gate: gate_state.map(|(_, base)| GateReport {
2099            baseline: base,
2100            evals: gate_evals,
2101            chosen: gate_chosen,
2102        }),
2103    };
2104    save_polished(model, out, &fm, &ts, o1, hp, &report)?;
2105    Ok(report)
2106}
2107
2108/// Hot-swap the trainable LN/FFN master copies into a runtime Pipeline
2109/// (frozen tensors stay mmap-backed — this reproduces the artifact the
2110/// polish would write, without writing it).
2111fn apply_trainables(pipe: &mut Pipeline, fm: &FcdModel, ts: &TrainState) {
2112    let hidden = fm.hidden;
2113    for (slot, &li) in ts.layers.iter().enumerate() {
2114        let b = slot * PARAMS_PER_LAYER;
2115        let inter = fm.layers[li].inter;
2116        let lw = &mut pipe.weights.layers[li];
2117        lw.input_norm = ts.data[b].clone();
2118        lw.post_norm = ts.data[b + 1].clone();
2119        lw.ffn = FfnKind::Dense(DenseFfn {
2120            gate_proj: QTensor::from_f32(ts.data[b + 2].clone(), inter, hidden),
2121            up_proj: QTensor::from_f32(ts.data[b + 3].clone(), inter, hidden),
2122            down_proj: QTensor::from_f32(ts.data[b + 4].clone(), hidden, inter),
2123            act: crate::pipeline::Act::Silu,
2124        });
2125    }
2126}
2127
2128/// Greedy loop-score probe through the streaming runtime.
2129fn gate_gen_scores(pipe: &mut Pipeline, g: &GenGateCfg) -> Result<Vec<f64>, String> {
2130    g.prompts
2131        .iter()
2132        .map(|p| {
2133            pipe.generate_from_ids(p, g.gen_tokens, None, None)
2134                .map(|r| loop_score(&r.token_ids))
2135        })
2136        .collect()
2137}
2138
2139/// Write the polished container: every source tensor byte-copied except
2140/// the converted layers' LN/FFN, which become f32 (per-tensor dtypes
2141/// are first-class in the directory — no requant noise on fresh
2142/// weights). Adds `provenance.o1_attn` + `provenance.fcd`.
2143fn save_polished(
2144    model: &CmfModel,
2145    out: &std::path::Path,
2146    fm: &FcdModel,
2147    ts: &TrainState,
2148    o1: &O1Cfg,
2149    hp: &FcdHyper,
2150    report: &FcdReport,
2151) -> Result<(), String> {
2152    use cortiq_core::format::TensorSpec;
2153    let mut replace: std::collections::HashMap<String, (usize, usize)> =
2154        std::collections::HashMap::new(); // name → (slot, param idx)
2155    for (s, &li) in ts.layers.iter().enumerate() {
2156        let p = format!("model.layers.{li}.");
2157        for (k, suffix) in [
2158            (0usize, "input_layernorm.weight"),
2159            (1, "post_attention_layernorm.weight"),
2160            (2, "mlp.gate_proj.weight"),
2161            (3, "mlp.up_proj.weight"),
2162            (4, "mlp.down_proj.weight"),
2163        ] {
2164            replace.insert(format!("{p}{suffix}"), (s, k));
2165        }
2166    }
2167    let mut specs = Vec::with_capacity(model.tensors.len());
2168    for t in &model.tensors {
2169        if let Some(&(s, k)) = replace.get(&t.name) {
2170            let data = &ts.data[s * PARAMS_PER_LAYER + k];
2171            let mut bytes = Vec::with_capacity(data.len() * 4);
2172            for v in data {
2173                bytes.extend_from_slice(&v.to_le_bytes());
2174            }
2175            specs.push(TensorSpec {
2176                name: t.name.clone(),
2177                dtype: TensorDtype::F32,
2178                shape: t.shape.clone(),
2179                data: bytes,
2180            });
2181        } else {
2182            specs.push(TensorSpec {
2183                name: t.name.clone(),
2184                dtype: t.dtype,
2185                shape: t.shape.clone(),
2186                data: model.entry_bytes(t).to_vec(),
2187            });
2188        }
2189    }
2190
2191    let mut header = model.header.clone();
2192    let mut prov = match header.provenance.take() {
2193        Some(serde_json::Value::Object(m)) => m,
2194        _ => serde_json::Map::new(),
2195    };
2196    let layers_json = match &o1.layers {
2197        O1Layers::All => serde_json::json!("all"),
2198        O1Layers::Deep(n) => serde_json::json!(format!("deep{n}")),
2199        O1Layers::List(v) => serde_json::json!(v),
2200    };
2201    prov.insert(
2202        "o1_attn".into(),
2203        serde_json::json!({
2204            "layers": layers_json, "m": o1.m, "w": o1.w, "sink": o1.sink
2205        }),
2206    );
2207    prov.insert(
2208        "fcd".into(),
2209        serde_json::json!({
2210            "steps": hp.steps, "lr": hp.lr, "kl_w": hp.kl_w,
2211            "bs": hp.bs, "seq": hp.seq,
2212            "teacher_ppl": report.teacher_ppl,
2213            "ppl_start": report.ppl_start,
2214            "ppl_final": report.ppl_final,
2215            "best_step": report.best_step,
2216            "converted_layers": report.converted,
2217        }),
2218    );
2219    header.provenance = Some(serde_json::Value::Object(prov));
2220    let _ = fm; // geometry only used for validation today
2221
2222    let masks = if model.masks.masks.is_empty() {
2223        None
2224    } else {
2225        Some(&model.masks)
2226    };
2227    CmfModel::write(out, &header, &specs, masks, model.vocab.as_deref())
2228        .map_err(|e| format!("writing polished cmf: {e}"))
2229}
2230
2231#[cfg(test)]
2232mod tests {
2233    use super::*;
2234
2235    /// Claim-13 selection: lowest ppl AMONG PASSING, not global lowest.
2236    #[test]
2237    fn gate_selects_lowest_ppl_among_passing() {
2238        let base = vec![0.10, 0.00, 0.20];
2239        let evals = vec![
2240            (25usize, 21.0, vec![0.10, 0.05, 0.20]), // pass
2241            (50, 18.0, vec![0.40, 0.00, 0.10]),      // fail: 0.40 > threshold
2242            (75, 19.0, vec![0.15, 0.05, 0.25]),      // pass — best passing
2243            (100, 18.5, vec![0.20, 0.30, 0.20]),     // fail: 0.30 > base+0.10
2244        ];
2245        let sel = select_checkpoint(&evals, &base, 0.35, 0.10);
2246        assert_eq!(sel, Some(2), "step 75 is the lowest-ppl PASSING checkpoint");
2247    }
2248
2249    /// All checkpoints fail → identity (None): the polish must never
2250    /// make generation worse than conversion alone.
2251    #[test]
2252    fn gate_all_fail_is_identity() {
2253        let base = vec![0.0, 0.0, 0.0];
2254        let evals = vec![
2255            (25usize, 15.0, vec![0.50, 0.0, 0.0]),
2256            (50, 14.0, vec![0.0, 0.36, 0.0]),
2257            (75, 13.0, vec![0.0, 0.0, 0.11]), // 0.11 > 0 + 0.10 slack
2258        ];
2259        assert_eq!(select_checkpoint(&evals, &base, 0.35, 0.10), None);
2260    }
2261
2262    /// Boundary discipline: scores AT the threshold / AT base+slack pass
2263    /// ("exceeds" is strict); ties in ppl resolve to the earliest step.
2264    #[test]
2265    fn gate_boundaries_and_tie_break() {
2266        let base = vec![0.25];
2267        assert!(gate_pass(&[0.35], &base, 0.35, 0.10), "== threshold passes");
2268        assert!(
2269            gate_pass(&[0.35], &[0.25], 0.35, 0.10),
2270            "== base+slack passes"
2271        );
2272        assert!(!gate_pass(&[0.351], &base, 0.35, 0.10));
2273        assert!(!gate_pass(&[0.30], &[0.10], 0.35, 0.10), "0.30 > 0.10+0.10");
2274        let evals = vec![(25usize, 20.0, vec![0.10]), (50, 20.0, vec![0.10])];
2275        assert_eq!(
2276            select_checkpoint(&evals, &base, 0.35, 0.10),
2277            Some(0),
2278            "equal ppl → earliest checkpoint"
2279        );
2280    }
2281}