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