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