Skip to main content

cortiq_engine/
skillbake.rs

1//! Native DTG-MA skill bake (Patent 2) — no Python, no torch.
2//!
3//! The certified recipe of `converter/make_skill_l1fcd.py`, in Rust on
4//! the `FcdModel` f32 replica:
5//!
6//! - **Phase A** — a trainable L1 mask over FFN neurons (one logit per
7//!   neuron, applied to the input of down_proj as σ(m)): pure LM loss
8//!   on the task corpus + a progressive L1 penalty. Every 30 steps the
9//!   binarized mask (σ>τ) is scored on held-out chunks; the best
10//!   checkpoint — the *denoising bottom* — is restored at the end.
11//!   Pruning noise neurons IMPROVES the model before it starts to hurt.
12//! - **Phase B** — FCD: the FFN of the last N layers trains against the
13//!   same LM loss with the hard mask active (cosine LR), held-out
14//!   gated, best checkpoint restored.
15//!
16//! Attention (softmax and GDN alike) is FROZEN and carries no gradient
17//! — exactly like the reference recipe (`torch.no_grad()` around the
18//! attention branch): the backward walks the residual stream through
19//! the FFN chain only, which is what makes a pure-Rust backward small.
20
21use crate::fcd::{FcdModel, LnFfn};
22use crate::fcd_ops as ops;
23use crate::sampler::SplitMix64;
24use cortiq_core::CmfModel;
25use std::sync::Arc;
26
27/// Hyper-parameters — defaults are the certified recipe.
28#[derive(Clone, Debug)]
29pub struct BakeHyper {
30    pub steps_a: usize,
31    pub steps_b: usize,
32    pub l1_init: f64,
33    pub l1_step: f64,
34    pub eval_every: usize,
35    pub lr_a: f64,
36    pub lr_b: f64,
37    pub tau: f32,
38    pub fcd_layers: usize,
39    /// Independent fixed-length records per optimizer step.  The loss is
40    /// normalized over all focused targets in the batch.
41    pub batch: usize,
42    /// Focused records per cached final-FFN optimizer step.  This is separate
43    /// from `batch` because Phase A stores every layer's activations while a
44    /// one-layer FCD cache stores only two hidden vectors per record.
45    pub fcd_batch: usize,
46    pub seed: u64,
47    /// Target sparsity (0..1). When >0, the best checkpoint must have
48    /// at least this fraction of neurons pruned; if none qualifies the
49    /// highest-sparsity checkpoint is used.
50    pub target_sparsity: f64,
51    /// L1 aggression multiplier: scales both l1_init and l1_step.
52    /// >1.0 = harder pruning push, <1.0 = softer.
53    pub l1_mult: f64,
54    /// Effective unlooped mask logit at step zero. The per-visit value is
55    /// solved so that the product over loop visits equals sigmoid(init).
56    /// 2.0 preserves the native recipe; 4.0 reproduces the older DTG-MA
57    /// trading notebooks' near-identity start.
58    pub mask_init: f32,
59    /// Penalize softplus(logit), whose derivative is sigmoid(logit), instead
60    /// of penalizing sigmoid(logit) itself. This reproduces the older DTG-MA
61    /// recipe and avoids an extra (1-sigmoid) attenuation near an open gate.
62    pub softplus_l1: bool,
63    /// Select Phase-A checkpoints by held-out hard balanced accuracy when
64    /// focused class tokens are configured. Otherwise held-out PPL remains
65    /// the checkpoint metric.
66    pub checkpoint_accuracy: bool,
67    /// Optional strict lower bound for a focused checkpoint's raw accuracy.
68    /// A checkpoint is eligible only when its measured accuracy is greater
69    /// than this value. This lets callers impose a natural-distribution
70    /// majority guard instead of selecting from a balanced holdout.
71    pub checkpoint_min_accuracy: Option<f64>,
72    /// Optional strict lower bound for focused balanced accuracy. Combined
73    /// with `checkpoint_min_accuracy`, this prevents a majority-only mask
74    /// from becoming the shipped specialist.
75    pub checkpoint_min_balanced_accuracy: Option<f64>,
76    /// When a joint accuracy guard is configured, rank eligible checkpoints
77    /// by raw accuracy first, then balanced accuracy and PPL. The historical
78    /// balanced-first selector remains the default for compatibility.
79    pub checkpoint_raw_priority: bool,
80    /// Round each layer's kept-neuron count UP to a multiple of this
81    /// (0/1 = off). 32 keeps the defragged FFN on grouped codecs
82    /// (in % 32 == 0) and SIMD kernels off their scalar tails.
83    pub align: usize,
84    /// Force one FFN width across all layers (the max aligned count) —
85    /// the whole-token GPU graphs require a uniform intermediate size.
86    pub uniform_inter: bool,
87    /// When non-empty, LM loss is accumulated only where the next token
88    /// is one of these ids. The whole chunk is still forwarded as context.
89    /// This is useful for supervised corpora with a long input and a
90    /// one-token answer, where ordinary all-token LM loss would drown the
91    /// task signal in prompt reconstruction.
92    pub focus_tokens: Vec<u32>,
93    /// Optional token(s) that must immediately follow a focused target.
94    /// Supervised ChatML uses the one-token label followed by `<|im_end|>`;
95    /// this prevents label names mentioned inside the user instruction from
96    /// being mistaken for answer positions.
97    pub focus_follow_tokens: Vec<u32>,
98}
99
100impl Default for BakeHyper {
101    fn default() -> Self {
102        Self {
103            steps_a: 240,
104            steps_b: 120,
105            l1_init: 0.01,
106            l1_step: 0.005,
107            eval_every: 30,
108            lr_a: 0.1,
109            lr_b: 1e-5,
110            tau: 0.5,
111            fcd_layers: 4,
112            batch: 1,
113            fcd_batch: 128,
114            seed: 0,
115            target_sparsity: 0.0,
116            l1_mult: 1.0,
117            mask_init: 2.0,
118            softplus_l1: false,
119            checkpoint_accuracy: false,
120            checkpoint_min_accuracy: None,
121            checkpoint_min_balanced_accuracy: None,
122            checkpoint_raw_priority: false,
123            align: 32,
124            uniform_inter: false,
125            focus_tokens: Vec::new(),
126            focus_follow_tokens: Vec::new(),
127        }
128    }
129}
130
131/// What the bake measured and produced.
132pub struct BakeReport {
133    /// Held-out PPL of the untouched backbone.
134    pub backbone: f64,
135    /// Held-out PPL with the best hard mask (the denoising bottom).
136    pub masked: f64,
137    /// Held-out PPL after FCD (the final specialist).
138    pub overlaid: f64,
139    pub pruned_ratio: f64,
140    pub kept_per_layer: Vec<usize>,
141    /// Hard focused-label accuracy, present when focus tokens were supplied.
142    pub backbone_accuracy: Option<f64>,
143    pub masked_accuracy: Option<f64>,
144    pub overlaid_accuracy: Option<f64>,
145    /// Macro recall over focused labels. Unlike raw accuracy this cannot be
146    /// improved by collapsing to the majority UP/DOWN class.
147    pub backbone_balanced_accuracy: Option<f64>,
148    pub masked_balanced_accuracy: Option<f64>,
149    pub overlaid_balanced_accuracy: Option<f64>,
150    pub selected_step: usize,
151    pub sec: f64,
152}
153
154pub struct BakeCheckpoint {
155    pub step: usize,
156    pub l1: f64,
157    pub ppl: f64,
158    pub sparsity: f64,
159    pub accuracy: Option<f64>,
160    pub balanced_accuracy: Option<f64>,
161}
162
163/// The trained artifacts: everything the defrag writer needs, f32.
164pub struct BakeArtifacts {
165    /// Per-PHYSICAL-layer live flags: the union over visits — a weight
166    /// row is removable from disk only when no visit keeps it.
167    pub keep: Vec<Vec<bool>>,
168    /// Per-VIRTUAL-layer live flags (physical × loops, pass-major): the
169    /// mask the file ships and the runtime applies per visit.
170    pub keep_visits: Vec<Vec<bool>>,
171    /// Per-layer down_proj `[hidden, inter]` with dead columns zeroed
172    /// (FCD layers: the trained weights; others: the backbone's).
173    pub down: Vec<Vec<f32>>,
174    /// Trained gate/up for the FCD layers (`None` elsewhere).
175    pub gate_up: Vec<Option<(Vec<f32>, Vec<f32>)>>,
176    /// Which layers went through Phase B.
177    pub fcd_layers: Vec<usize>,
178    /// The trained mask logits, per virtual layer — a CONTINUOUS
179    /// per-neuron importance the hard keep flags throw away. The tube
180    /// planner ranks and orders neurons by these, not by raw
181    /// activation mass.
182    pub logits: Vec<Vec<f32>>,
183    /// Phase-A logits after the last requested optimization step, before
184    /// restoring the selected hard-validation checkpoint.
185    pub final_logits: Vec<Vec<f32>>,
186    pub checkpoints: Vec<BakeCheckpoint>,
187}
188
189const CLIP: f64 = 1.0;
190const B1: f64 = 0.9;
191const B2: f64 = 0.999;
192const EPS: f64 = 1e-8;
193
194/// Plain Adam over a set of f32 tensors (masks are tiny, FFN mid-size).
195struct Adam {
196    m: Vec<Vec<f64>>,
197    v: Vec<Vec<f64>>,
198    t: i32,
199    lr: f64,
200}
201
202impl Adam {
203    fn new(sizes: &[usize], lr: f64) -> Self {
204        Self {
205            m: sizes.iter().map(|&n| vec![0.0; n]).collect(),
206            v: sizes.iter().map(|&n| vec![0.0; n]).collect(),
207            t: 0,
208            lr,
209        }
210    }
211
212    /// Global-norm clip + Adam step. `params[i].len() == grads[i].len()`.
213    fn step(&mut self, params: &mut [&mut [f32]], grads: &[Vec<f64>], lr_scale: f64) {
214        let gn: f64 = grads
215            .iter()
216            .flat_map(|g| g.iter().map(|x| x * x))
217            .sum::<f64>()
218            .sqrt();
219        let clip = if gn > CLIP { CLIP / gn } else { 1.0 };
220        self.t += 1;
221        let (bc1, bc2) = (1.0 - B1.powi(self.t), 1.0 - B2.powi(self.t));
222        for (pi, p) in params.iter_mut().enumerate() {
223            for j in 0..p.len() {
224                let g = grads[pi][j] * clip;
225                let m = &mut self.m[pi][j];
226                let v = &mut self.v[pi][j];
227                *m = B1 * *m + (1.0 - B1) * g;
228                *v = B2 * *v + (1.0 - B2) * g * g;
229                let upd = (*m / bc1) / ((*v / bc2).sqrt() + EPS);
230                p[j] -= (self.lr * lr_scale * upd) as f32;
231            }
232        }
233    }
234}
235
236/// Mask logit at step zero, solved for the loop depth.
237///
238/// The gate multiplies the FFN once per VISIT, so a Looped Transformer
239/// applies it `loops` times per token and the factor compounds. What
240/// must be held constant across depths is the EFFECTIVE start — the
241/// product the stack actually sees — at the value the recipe was
242/// validated with on ordinary models, σ(2.0) = 0.881:
243///
244/// ```text
245/// σ(m0)^loops = σ(2.0)   →   m0 = logit( σ(2.0)^(1/loops) )
246/// ```
247///
248/// `loops = 1` returns 2.0 exactly, so nothing regresses. Two known
249/// wrong answers this replaces: the old hardcoded 2.0, which at two
250/// visits compounds to 0.776 and took Nanbeige 4.2 from a baseline of
251/// 4.187 to 278.4 at step 30; and a start pushed to identity, which
252/// cannot learn because the update carries σ'(m) = σ(1−σ), worth 5e-4
253/// at σ = 0.9995 against 0.105 at 2.0.
254pub fn mask_init_logit_for(loops: usize, effective_logit: f32) -> f32 {
255    let base = 1.0f32 / (1.0 + (-effective_logit).exp());
256    let per_visit = base.powf(1.0 / loops.max(1) as f32);
257    (per_visit / (1.0 - per_visit)).ln()
258}
259
260pub fn mask_init_logit(loops: usize) -> f32 {
261    mask_init_logit_for(loops, 2.0)
262}
263
264/// Learning-rate scale for the mask step, given the loop depth.
265///
266/// The backward accumulates every visit of a physical layer into the
267/// same mask gradient, so an unscaled step is `loops` times the tuned
268/// one. One step should mean one token's worth of movement at any depth.
269pub fn mask_step_scale(loops: usize) -> f64 {
270    1.0 / loops.max(1) as f64
271}
272
273fn sigmoid(x: f32) -> f32 {
274    1.0 / (1.0 + (-x).exp())
275}
276
277fn sparsity_grad(logit: f32, softplus_l1: bool) -> f64 {
278    let s = sigmoid(logit) as f64;
279    if softplus_l1 { s } else { s * (1.0 - s) }
280}
281
282fn is_scored_target(
283    ids: &[u32],
284    target_index: usize,
285    sequence_end: usize,
286    focus: &[u32],
287    follow: &[u32],
288) -> bool {
289    if focus.is_empty() {
290        return true;
291    }
292    focus.contains(&ids[target_index])
293        && (follow.is_empty()
294            || (target_index + 1 < sequence_end && follow.contains(&ids[target_index + 1])))
295}
296
297/// One forward + CE(+optionally backward through the FFN chain).
298/// Returns (nll_sum, tokens). `dmask`/`dffn` accumulate when given.
299struct Pass<'a> {
300    fm: &'a FcdModel,
301    tau: f32,
302    /// σ(m) per layer when soft; binarized when `hard`.
303    logits: &'a [Vec<f32>],
304    hard: bool,
305    /// Phase-B replacement FFN weights per layer (trained copies).
306    ffn: &'a [Option<(Vec<f32>, Vec<f32>, Vec<f32>)>],
307    /// Empty means ordinary all-token LM loss.
308    focus_tokens: &'a [u32],
309    /// Empty means no right-context constraint on focused targets.
310    focus_follow_tokens: &'a [u32],
311}
312
313#[derive(Clone, Debug, Default)]
314struct FocusStats {
315    total: usize,
316    correct: usize,
317    class_total: Vec<usize>,
318    class_correct: Vec<usize>,
319}
320
321impl FocusStats {
322    fn new(classes: usize) -> Self {
323        Self {
324            class_total: vec![0; classes],
325            class_correct: vec![0; classes],
326            ..Self::default()
327        }
328    }
329
330    fn accuracy(&self) -> Option<f64> {
331        (self.total > 0).then(|| self.correct as f64 / self.total as f64)
332    }
333
334    fn balanced_accuracy(&self) -> Option<f64> {
335        let recalls: Vec<f64> = self
336            .class_total
337            .iter()
338            .zip(&self.class_correct)
339            .filter_map(|(&n, &ok)| (n > 0).then(|| ok as f64 / n as f64))
340            .collect();
341        (!recalls.is_empty()).then(|| recalls.iter().sum::<f64>() / recalls.len() as f64)
342    }
343
344    fn merge(&mut self, other: &Self) {
345        self.total += other.total;
346        self.correct += other.correct;
347        if self.class_total.len() < other.class_total.len() {
348            self.class_total.resize(other.class_total.len(), 0);
349            self.class_correct.resize(other.class_correct.len(), 0);
350        }
351        for (dst, src) in self.class_total.iter_mut().zip(&other.class_total) {
352            *dst += src;
353        }
354        for (dst, src) in self.class_correct.iter_mut().zip(&other.class_correct) {
355            *dst += src;
356        }
357    }
358}
359
360#[derive(Clone, Debug)]
361struct HeldScore {
362    ppl: f64,
363    accuracy: Option<f64>,
364    balanced_accuracy: Option<f64>,
365}
366
367/// Frozen boundary immediately before one trainable final FFN.  This is not
368/// an adapter or a donor checkpoint: both vectors are extracted directly from
369/// the opened CMF, kept in RAM, and discarded when the native bake ends.
370#[derive(Default)]
371struct FocusedFcdCache {
372    h1: Vec<f32>,
373    n2: Vec<f32>,
374    targets: Vec<usize>,
375}
376
377impl FocusedFcdCache {
378    fn len(&self) -> usize {
379        self.targets.len()
380    }
381
382    fn append(&mut self, mut other: Self) {
383        self.h1.append(&mut other.h1);
384        self.n2.append(&mut other.n2);
385        self.targets.append(&mut other.targets);
386    }
387}
388
389fn gather_rows(values: &[f32], rows: &[usize], width: usize) -> Vec<f32> {
390    let mut out = Vec::with_capacity(rows.len() * width);
391    for &row in rows {
392        out.extend_from_slice(&values[row * width..(row + 1) * width]);
393    }
394    out
395}
396
397impl Pass<'_> {
398    fn gates(&self, li: usize) -> Vec<f32> {
399        self.logits[li]
400            .iter()
401            .map(|&l| {
402                let s = sigmoid(l);
403                if self.hard {
404                    if s > self.tau { 1.0 } else { 0.0 }
405                } else {
406                    s
407                }
408            })
409            .collect()
410    }
411
412    fn wts<'b>(&'b self, li: usize, mats: &'b crate::fcd::LayerMats) -> LnFfn<'b> {
413        let l = &self.fm.layers[li];
414        match &self.ffn[li] {
415            Some((g, u, d)) => LnFfn {
416                iln: &l.iln,
417                pln: &l.pln,
418                gate: g,
419                up: u,
420                down: d,
421                // Trained copies move every Adam step — no prebuilt concat.
422                gu: None,
423            },
424            None => LnFfn {
425                iln: &l.iln,
426                pln: &l.pln,
427                gate: &[],
428                up: &[],
429                down: &mats.down,
430                gu: Some(&mats.gu),
431            },
432        }
433    }
434
435    /// Extract the exact frozen boundary of the final FFN for focused answer
436    /// positions.  Only the ordinary one-pass/final-layer case is cacheable:
437    /// looped stacks revisit the same FFN after its own changed output and
438    /// correctly fall back to the full Phase-B path.
439    fn cache_final_ffn_batch(&self, ids: &[u32], batch: usize) -> Result<FocusedFcdCache, String> {
440        let fm = self.fm;
441        if fm.loops.max(1) != 1 || fm.layers.is_empty() || self.focus_tokens.is_empty() {
442            return Err("focused final-FFN cache needs a one-pass stack and focus tokens".into());
443        }
444        if ids.len() % batch.max(1) != 0 {
445            return Err("focused final-FFN cache received a ragged batch".into());
446        }
447        let t = ids.len() / batch.max(1);
448        let hsz = fm.hidden;
449        let last = fm.layers.len() - 1;
450        let mut sources = Vec::new();
451        let mut targets = Vec::new();
452        for bi in 0..batch {
453            let base = bi * t;
454            for target_index in base + 1..base + t {
455                if !is_scored_target(
456                    ids,
457                    target_index,
458                    base + t,
459                    self.focus_tokens,
460                    self.focus_follow_tokens,
461                ) {
462                    continue;
463                }
464                sources.push(target_index - 1);
465                targets.push(
466                    self.focus_tokens
467                        .iter()
468                        .position(|&id| id == ids[target_index])
469                        .expect("focused target belongs to focus_tokens"),
470                );
471            }
472        }
473        if sources.is_empty() {
474            return Ok(FocusedFcdCache::default());
475        }
476
477        let mut hidden = vec![0f32; ids.len() * hsz];
478        for (row, &id) in ids.iter().enumerate() {
479            hidden[row * hsz..(row + 1) * hsz]
480                .copy_from_slice(&fm.embed[id as usize * hsz..(id as usize + 1) * hsz]);
481        }
482        for layer in 0..=last {
483            let gate = self.gates(layer);
484            let mats = fm.mats(layer)?;
485            let weights = self.wts(layer, &mats);
486            let (next, acts) = fm.layer_forward_scaled(
487                layer,
488                &hidden,
489                batch,
490                t,
491                &weights,
492                false,
493                layer == last,
494                Some(&gate),
495            );
496            if layer == last {
497                let acts = acts.expect("last FFN boundary requested");
498                return Ok(FocusedFcdCache {
499                    h1: gather_rows(&acts.h1, &sources, hsz),
500                    n2: gather_rows(&acts.n2, &sources, hsz),
501                    targets,
502                });
503            }
504            hidden = next;
505        }
506        unreachable!("non-empty stack has a final layer")
507    }
508
509    /// Teacher-forced NLL over one chunk; when `grad` is set, backprop
510    /// through the FFN chain into the mask grads (and FFN grads for
511    /// Phase-B layers).
512    #[allow(clippy::too_many_arguments)]
513    fn chunk(
514        &self,
515        ids: &[u32],
516        grad: Option<(
517            &mut [Vec<f64>],
518            &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
519        )>,
520    ) -> (f64, usize) {
521        self.chunk_batch(ids, 1, grad)
522    }
523
524    /// `chunk` over `b` equal-length sequences flattened into `ids`.
525    ///
526    /// Everything under this level was batch-aware all along
527    /// (`layer_forward_scaled` and every attention fwd/bwd take `b`);
528    /// only this wrapper hardcoded 1. Evaluation is where it pays: the
529    /// held set is 12 chunks scored one at a time, which on a 4 B model
530    /// meant 12× the GEMM submits for the same arithmetic.
531    fn chunk_batch(
532        &self,
533        ids: &[u32],
534        b: usize,
535        grad: Option<(
536            &mut [Vec<f64>],
537            &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
538        )>,
539    ) -> (f64, usize) {
540        self.chunk_batch_scored(ids, b, grad, None)
541    }
542
543    fn chunk_batch_scored(
544        &self,
545        ids: &[u32],
546        b: usize,
547        grad: Option<(
548            &mut [Vec<f64>],
549            &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
550        )>,
551        mut focus_stats: Option<&mut FocusStats>,
552    ) -> (f64, usize) {
553        let fm = self.fm;
554        let hsz = fm.hidden;
555        debug_assert!(ids.len() % b.max(1) == 0, "ragged batch");
556        let t = ids.len() / b.max(1);
557        let n = b * t;
558        let nl = fm.layers.len();
559        // Embed.
560        let mut h = vec![0f32; n * hsz];
561        for (r, &id) in ids.iter().enumerate() {
562            h[r * hsz..(r + 1) * hsz]
563                .copy_from_slice(&fm.embed[id as usize * hsz..(id as usize + 1) * hsz]);
564        }
565        // Forward over VIRTUAL layers: a Looped Transformer runs the
566        // stack `fm.loops` times, with a final_norm at each loop
567        // boundary when the file says so. Everything below indexes
568        // activations by the virtual step and weights/grads by the
569        // physical layer `vl % nl` — so a physical layer visited twice
570        // accumulates both visits' gradients, which is what the loop
571        // means mathematically.
572        let loops = fm.loops.max(1);
573        let vn = nl * loops;
574        let mut h_ins = Vec::with_capacity(vn);
575        let mut acts = Vec::with_capacity(vn);
576        let mut masks = Vec::with_capacity(vn);
577        // Loop-boundary norms, saved for the backward: (input, inv).
578        let mut lnorms: Vec<Option<(Vec<f32>, Vec<f32>)>> = vec![None; vn];
579        for vl in 0..vn {
580            let li = vl % nl;
581            // The gate is PER VISIT: the two passes of a loop are
582            // different computations sharing one set of weights, so the
583            // mask must be allowed to differ between them. Weights stay
584            // indexed by the physical layer.
585            let g = self.gates(vl);
586            let mats_hold = fm.mats(li).expect("layer mats");
587            let wts = self.wts(li, &mats_hold);
588            let want = grad.is_some();
589            let (h2, a) = fm.layer_forward_scaled(li, &h, b, t, &wts, false, want, Some(&g));
590            h_ins.push(if want { h } else { Vec::new() });
591            acts.push(a);
592            masks.push(g);
593            h = h2;
594            // Mid-stack norm at every loop boundary except the last —
595            // the final one folds into the head below.
596            if fm.loop_norm && li + 1 == nl && vl + 1 < vn {
597                let mut hn = vec![0f32; n * hsz];
598                let mut inv = vec![0f32; n];
599                ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
600                if want {
601                    lnorms[vl] = Some((h, inv));
602                }
603                h = hn;
604            }
605        }
606        // Final norm + tied LM head, CE summed over positions 1..t.
607        let mut hn = vec![0f32; n * hsz];
608        let mut inv = vec![0f32; n];
609        ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
610        let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
611        let vocab = lm.len() / hsz;
612        let pool = fm.pool.as_deref();
613        let mut nll = 0f64;
614        let mut dh_n = vec![0f32; n * hsz]; // dL/d hn
615        // Chunk the vocab matmul over positions to bound the logits buf.
616        // Positions are walked PER SEQUENCE: the last position of chunk
617        // i must not be scored against the first token of chunk i+1.
618        const POS_CHUNK: usize = 64;
619        let scored = (0..b)
620            .map(|bi| {
621                let base = bi * t;
622                (base + 1..base + t)
623                    .filter(|&target_index| {
624                        is_scored_target(
625                            ids,
626                            target_index,
627                            base + t,
628                            self.focus_tokens,
629                            self.focus_follow_tokens,
630                        )
631                    })
632                    .count()
633            })
634            .sum::<usize>();
635        if scored == 0 {
636            return (0.0, 0);
637        }
638        if self.focus_tokens.is_empty() {
639            // Ordinary language-model mode still needs the complete
640            // vocabulary distribution at every target position.
641            for bi in 0..b {
642                let base = bi * t;
643                let mut p0 = 0usize;
644                while p0 < t - 1 {
645                    let pc = POS_CHUNK.min(t - 1 - p0);
646                    let mut logits = vec![0f32; pc * vocab];
647                    ops::gemm_nt(
648                        &hn[(base + p0) * hsz..(base + p0 + pc) * hsz],
649                        lm,
650                        &mut logits,
651                        pc,
652                        hsz,
653                        vocab,
654                        pool,
655                    );
656                    for r in 0..pc {
657                        let target_index = base + p0 + r + 1;
658                        let target = ids[target_index] as usize;
659                        let row = &mut logits[r * vocab..(r + 1) * vocab];
660                        let mx = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
661                        let mut sum = 0f64;
662                        for v in row.iter() {
663                            sum += ((*v as f64) - mx).exp();
664                        }
665                        nll += mx + sum.ln() - row[target] as f64;
666                        if grad.is_some() {
667                            // dCE/dlogit = softmax − onehot, scaled by 1/scored.
668                            let inv_n = 1.0 / scored as f64;
669                            for v in row.iter_mut() {
670                                *v = ((((*v as f64) - mx).exp() / sum) * inv_n) as f32;
671                            }
672                            row[target] -= inv_n as f32;
673                        }
674                    }
675                    if grad.is_some() {
676                        ops::gemm_dx(
677                            &logits,
678                            lm,
679                            &mut dh_n[(base + p0) * hsz..(base + p0 + pc) * hsz],
680                            pc,
681                            hsz,
682                            vocab,
683                            pool,
684                        );
685                    }
686                    p0 += pc;
687                }
688            }
689        } else {
690            // Exact classifier mode. Only the declared label rows can enter
691            // either the normalizer or its gradient, so multiplying every
692            // hidden by the full ~250k-row LM head is pure wasted work. It
693            // made a two-label archaeology run spend more time in the head
694            // than in mask Adam and limited experiments to a tiny fraction
695            // of the updates used by the original notebooks.
696            let inv_n = 1.0 / scored as f64;
697            for bi in 0..b {
698                let base = bi * t;
699                for target_index in base + 1..base + t {
700                    if !is_scored_target(
701                        ids,
702                        target_index,
703                        base + t,
704                        self.focus_tokens,
705                        self.focus_follow_tokens,
706                    ) {
707                        continue;
708                    }
709                    let target_class = self
710                        .focus_tokens
711                        .iter()
712                        .position(|&id| id == ids[target_index])
713                        .expect("focused target belongs to focus_tokens");
714                    let source = target_index - 1;
715                    let hidden = &hn[source * hsz..(source + 1) * hsz];
716                    let class_logits: Vec<f32> = self
717                        .focus_tokens
718                        .iter()
719                        .map(|&id| {
720                            let row = &lm[id as usize * hsz..(id as usize + 1) * hsz];
721                            hidden.iter().zip(row).map(|(&x, &w)| x * w).sum()
722                        })
723                        .collect();
724                    let mx = class_logits
725                        .iter()
726                        .copied()
727                        .fold(f32::NEG_INFINITY, f32::max) as f64;
728                    let probs: Vec<f64> = class_logits
729                        .iter()
730                        .map(|&value| ((value as f64) - mx).exp())
731                        .collect();
732                    let sum: f64 = probs.iter().sum();
733                    nll += mx + sum.ln() - class_logits[target_class] as f64;
734                    if let Some(stats) = focus_stats.as_deref_mut() {
735                        let predicted_class = class_logits
736                            .iter()
737                            .enumerate()
738                            .max_by(|(_, left), (_, right)| left.total_cmp(right))
739                            .map(|(index, _)| index)
740                            .expect("focus_tokens is non-empty");
741                        stats.total += 1;
742                        stats.class_total[target_class] += 1;
743                        if predicted_class == target_class {
744                            stats.correct += 1;
745                            stats.class_correct[target_class] += 1;
746                        }
747                    }
748                    if grad.is_some() {
749                        let dh = &mut dh_n[source * hsz..(source + 1) * hsz];
750                        for (class, (&id, probability)) in
751                            self.focus_tokens.iter().zip(probs).enumerate()
752                        {
753                            let coefficient = (probability / sum
754                                - usize::from(class == target_class) as f64)
755                                * inv_n;
756                            let row = &lm[id as usize * hsz..(id as usize + 1) * hsz];
757                            for (value, &weight) in dh.iter_mut().zip(row) {
758                                *value += (coefficient * weight as f64) as f32;
759                            }
760                        }
761                    }
762                }
763            }
764        }
765        let Some((dmask, dffn)) = grad else {
766            return (nll, scored);
767        };
768        // Backward: final norm, then the FFN chain layer by layer.
769        let t_bwd = std::time::Instant::now();
770        let mut dh = vec![0f32; n * hsz];
771        ops::rmsnorm_bwd(&h, &fm.final_norm, &inv, &dh_n, fm.gemma, &mut dh, None);
772        for vl in (0..vn).rev() {
773            let li = vl % nl;
774            // Undo the loop-boundary norm this step fed into.
775            if let Some((hb, inv)) = lnorms[vl].as_ref() {
776                let mut dprev = vec![0f32; n * hsz];
777                ops::rmsnorm_bwd(hb, &fm.final_norm, inv, &dh, fm.gemma, &mut dprev, None);
778                dh = dprev;
779            }
780            let a = acts[vl].as_ref().expect("acts saved in grad mode");
781            let g = &masks[vl];
782            let inter = fm.layers[li].inter;
783            let mats_hold = fm.mats(li).expect("layer mats");
784            let wts = self.wts(li, &mats_hold);
785            // h2 = h1 + act2 @ downᵀ  →  dact2 = dh @ down.
786            let mut dact2 = vec![0f32; n * inter];
787            ops::gemm_dx(&dh, wts.down, &mut dact2, n, inter, hsz, fm.pool.as_deref());
788            if let Some((_, _, dd)) = dffn[li].as_mut() {
789                // dW_down += dhᵀ · act2 (act2 = act·g).
790                let mut act2 = a.act.clone();
791                for r in 0..n {
792                    for (x, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
793                        *x *= gv;
794                    }
795                }
796                let mut dw = vec![0f32; hsz * inter];
797                ops::gemm_dw(&dh, &act2, &mut dw, n, inter, hsz, fm.pool.as_deref());
798                for (o, &x) in dd.iter_mut().zip(&dw) {
799                    *o += x as f64;
800                }
801            }
802            // Mask grad: dm = Σ_t dact2·act · σ'(m)  (soft; STE-equal).
803            // Indexed by the VIRTUAL layer: each visit's mask row gets
804            // exactly its own visit's gradient, no cross-visit sum.
805            {
806                let dm = &mut dmask[vl];
807                for r in 0..n {
808                    let da = &dact2[r * inter..(r + 1) * inter];
809                    let aa = &a.act[r * inter..(r + 1) * inter];
810                    for j in 0..inter {
811                        dm[j] += da[j] as f64 * aa[j] as f64;
812                    }
813                }
814                // σ'(m) folded in once per chunk (constant per neuron).
815                for (j, d) in dm.iter_mut().enumerate() {
816                    let _ = j;
817                    let _ = d;
818                }
819            }
820            // dact = dact2 · g;  silu·mul backward.
821            let mut dg_pre = vec![0f32; n * inter];
822            let mut du_pre = vec![0f32; n * inter];
823            for r in 0..n {
824                for j in 0..inter {
825                    let i = r * inter + j;
826                    let da = dact2[i] * g[j];
827                    let sg = ops::silu(a.gpre[i]);
828                    dg_pre[i] = da * a.upre[i] * ops::silu_bwd(a.gpre[i]);
829                    du_pre[i] = da * sg;
830                }
831            }
832            // dn2 = dg_pre @ gate + du_pre @ up — one fused submit when
833            // the frozen concat exists; the trained-copy path keeps two.
834            let mut dn2 = vec![0f32; n * hsz];
835            if let Some(gu) = wts.gu {
836                let mut dgu = vec![0f32; n * 2 * inter];
837                for r in 0..n {
838                    let row = &mut dgu[r * 2 * inter..(r + 1) * 2 * inter];
839                    row[..inter].copy_from_slice(&dg_pre[r * inter..(r + 1) * inter]);
840                    row[inter..].copy_from_slice(&du_pre[r * inter..(r + 1) * inter]);
841                }
842                ops::gemm_dx(&dgu, gu, &mut dn2, n, hsz, 2 * inter, fm.pool.as_deref());
843            } else {
844                ops::gemm_dx(
845                    &dg_pre,
846                    wts.gate,
847                    &mut dn2,
848                    n,
849                    hsz,
850                    inter,
851                    fm.pool.as_deref(),
852                );
853                let mut dn2b = vec![0f32; n * hsz];
854                ops::gemm_dx(
855                    &du_pre,
856                    wts.up,
857                    &mut dn2b,
858                    n,
859                    hsz,
860                    inter,
861                    fm.pool.as_deref(),
862                );
863                for (x, &y) in dn2.iter_mut().zip(&dn2b) {
864                    *x += y;
865                }
866            }
867            if let Some((dgw, duw, _)) = dffn[li].as_mut() {
868                let mut dw = vec![0f32; inter * hsz];
869                ops::gemm_dw(&dg_pre, &a.n2, &mut dw, n, hsz, inter, fm.pool.as_deref());
870                for (o, &x) in dgw.iter_mut().zip(&dw) {
871                    *o += x as f64;
872                }
873                dw.fill(0.0);
874                ops::gemm_dw(&du_pre, &a.n2, &mut dw, n, hsz, inter, fm.pool.as_deref());
875                for (o, &x) in duw.iter_mut().zip(&dw) {
876                    *o += x as f64;
877                }
878            }
879            // Post-norm backward into h1; the attention branch carries
880            // no gradient (frozen), so dh1 flows straight to dh_in.
881            let mut dh1 = dh.clone(); // residual h2 = h1 + ffn
882            ops::rmsnorm_bwd(&a.h1, wts.pln, &a.inv2, &dn2, fm.gemma, &mut dh1, None);
883            dh = dh1;
884            let _ = &h_ins[vl];
885        }
886        crate::fcd::prof::add(&crate::fcd::prof::BWD, t_bwd);
887        (nll, scored)
888    }
889}
890
891/// Held-out PPL with the hard mask (and Phase-B weights when present).
892fn held_ppl(pass: &Pass, held: &[Vec<u32>]) -> f64 {
893    held_score(pass, held).ppl
894}
895
896fn held_score(pass: &Pass, held: &[Vec<u32>]) -> HeldScore {
897    // Bounded batched passes over the held set: equal-length records share
898    // one GEMM per weight, while the cap keeps activation/GPU scratch memory
899    // predictable for a full validation sweep. Aggregation is exact because
900    // NLL and focused-class counts are additive across groups.
901    if held.is_empty() {
902        return HeldScore {
903            ppl: f64::NAN,
904            accuracy: None,
905            balanced_accuracy: None,
906        };
907    }
908    let mut nll = 0f64;
909    let mut n = 0usize;
910    let mut stats = FocusStats::new(pass.focus_tokens.len());
911    const GROUP: usize = 32;
912    for group in held.chunks(GROUP) {
913        let t = group[0].len();
914        if group.iter().all(|c| c.len() == t) {
915            let flat: Vec<u32> = group.iter().flatten().copied().collect();
916            let mut part = FocusStats::new(pass.focus_tokens.len());
917            let (l, k) = pass.chunk_batch_scored(&flat, group.len(), None, Some(&mut part));
918            nll += l;
919            n += k;
920            stats.merge(&part);
921        } else {
922            for c in group {
923                let mut part = FocusStats::new(pass.focus_tokens.len());
924                let (l, k) = pass.chunk_batch_scored(c, 1, None, Some(&mut part));
925                nll += l;
926                n += k;
927                stats.merge(&part);
928            }
929        }
930    }
931    HeldScore {
932        ppl: (nll / n.max(1) as f64).exp(),
933        accuracy: stats.accuracy(),
934        balanced_accuracy: stats.balanced_accuracy(),
935    }
936}
937
938fn calibration_batch(
939    calib: &[Vec<u32>],
940    step: usize,
941    requested: usize,
942) -> Result<(Vec<u32>, usize), String> {
943    let batch = requested.max(1).min(calib.len());
944    let width = calib[0].len();
945    let mut flat = Vec::with_capacity(batch * width);
946    for offset in 0..batch {
947        let record = &calib[(step * batch + offset) % calib.len()];
948        if record.len() != width {
949            return Err(format!(
950                "skill bake: --batch needs equal-length records ({} != {width})",
951                record.len()
952            ));
953        }
954        flat.extend_from_slice(record);
955    }
956    Ok((flat, batch))
957}
958
959fn build_focused_fcd_cache(
960    pass: &Pass<'_>,
961    records: &[Vec<u32>],
962    extraction_batch: usize,
963) -> Result<FocusedFcdCache, String> {
964    let mut cache = FocusedFcdCache::default();
965    for group in records.chunks(extraction_batch.max(1)) {
966        let width = group[0].len();
967        if group.iter().any(|record| record.len() != width) {
968            return Err("focused final-FFN cache needs equal-length records".into());
969        }
970        let flat: Vec<u32> = group.iter().flatten().copied().collect();
971        cache.append(pass.cache_final_ffn_batch(&flat, group.len())?);
972    }
973    Ok(cache)
974}
975
976/// Exact final-FFN focused loss, optionally with gradients for the three full
977/// FCD projections.  The upstream representation came from the same CMF and
978/// mask in `build_focused_fcd_cache`; no external checkpoint format is used.
979fn cached_fcd_run(
980    fm: &FcdModel,
981    cache: &FocusedFcdCache,
982    indices: &[usize],
983    weights: (&[f32], &[f32], &[f32]),
984    gate: &[f32],
985    focus_tokens: &[u32],
986    want_grad: bool,
987) -> (HeldScore, Option<(Vec<f64>, Vec<f64>, Vec<f64>)>) {
988    let (gate_w, up_w, down_w) = weights;
989    let rows = indices.len();
990    let hidden = fm.hidden;
991    let inter = gate.len();
992    if rows == 0 {
993        return (
994            HeldScore {
995                ppl: f64::NAN,
996                accuracy: None,
997                balanced_accuracy: None,
998            },
999            None,
1000        );
1001    }
1002    let h1 = gather_rows(&cache.h1, indices, hidden);
1003    let n2 = gather_rows(&cache.n2, indices, hidden);
1004    let mut gate_pre = vec![0f32; rows * inter];
1005    let mut up_pre = vec![0f32; rows * inter];
1006    ops::gemm_nt(
1007        &n2,
1008        gate_w,
1009        &mut gate_pre,
1010        rows,
1011        hidden,
1012        inter,
1013        fm.pool.as_deref(),
1014    );
1015    ops::gemm_nt(
1016        &n2,
1017        up_w,
1018        &mut up_pre,
1019        rows,
1020        hidden,
1021        inter,
1022        fm.pool.as_deref(),
1023    );
1024    let mut act = vec![0f32; rows * inter];
1025    for row in 0..rows {
1026        for column in 0..inter {
1027            let at = row * inter + column;
1028            act[at] = ops::silu(gate_pre[at]) * up_pre[at] * gate[column];
1029        }
1030    }
1031    let mut ffn = vec![0f32; rows * hidden];
1032    ops::gemm_nt(
1033        &act,
1034        down_w,
1035        &mut ffn,
1036        rows,
1037        inter,
1038        hidden,
1039        fm.pool.as_deref(),
1040    );
1041    let mut h2 = h1;
1042    for (value, &delta) in h2.iter_mut().zip(&ffn) {
1043        *value += delta;
1044    }
1045    let mut normed = vec![0f32; rows * hidden];
1046    let mut inv = vec![0f32; rows];
1047    ops::rmsnorm_fwd(&h2, &fm.final_norm, fm.eps, fm.gemma, &mut normed, &mut inv);
1048
1049    let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
1050    let mut stats = FocusStats::new(focus_tokens.len());
1051    let mut nll = 0f64;
1052    let mut dh_normed = want_grad.then(|| vec![0f32; rows * hidden]);
1053    for (local_row, &cache_row) in indices.iter().enumerate() {
1054        let target = cache.targets[cache_row];
1055        let state = &normed[local_row * hidden..(local_row + 1) * hidden];
1056        let logits: Vec<f32> = focus_tokens
1057            .iter()
1058            .map(|&id| {
1059                state
1060                    .iter()
1061                    .zip(&lm[id as usize * hidden..(id as usize + 1) * hidden])
1062                    .map(|(&left, &right)| left * right)
1063                    .sum()
1064            })
1065            .collect();
1066        let mx = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64;
1067        let exps: Vec<f64> = logits
1068            .iter()
1069            .map(|&value| ((value as f64) - mx).exp())
1070            .collect();
1071        let sum: f64 = exps.iter().sum();
1072        nll += mx + sum.ln() - logits[target] as f64;
1073        let predicted = logits
1074            .iter()
1075            .enumerate()
1076            .max_by(|(_, left), (_, right)| left.total_cmp(right))
1077            .map(|(class, _)| class)
1078            .expect("focused classes are non-empty");
1079        stats.total += 1;
1080        stats.class_total[target] += 1;
1081        if predicted == target {
1082            stats.correct += 1;
1083            stats.class_correct[target] += 1;
1084        }
1085        if let Some(gradient) = dh_normed.as_mut() {
1086            let row = &mut gradient[local_row * hidden..(local_row + 1) * hidden];
1087            for (class, (&id, probability)) in focus_tokens.iter().zip(exps).enumerate() {
1088                let coefficient =
1089                    (probability / sum - usize::from(class == target) as f64) / rows as f64;
1090                let head = &lm[id as usize * hidden..(id as usize + 1) * hidden];
1091                for (value, &weight) in row.iter_mut().zip(head) {
1092                    *value += (coefficient * weight as f64) as f32;
1093                }
1094            }
1095        }
1096    }
1097    let score = HeldScore {
1098        ppl: (nll / rows as f64).exp(),
1099        accuracy: stats.accuracy(),
1100        balanced_accuracy: stats.balanced_accuracy(),
1101    };
1102    let Some(dh_normed) = dh_normed else {
1103        return (score, None);
1104    };
1105
1106    let mut dh = vec![0f32; rows * hidden];
1107    ops::rmsnorm_bwd(
1108        &h2,
1109        &fm.final_norm,
1110        &inv,
1111        &dh_normed,
1112        fm.gemma,
1113        &mut dh,
1114        None,
1115    );
1116    let mut dact = vec![0f32; rows * inter];
1117    ops::gemm_dx(
1118        &dh,
1119        down_w,
1120        &mut dact,
1121        rows,
1122        inter,
1123        hidden,
1124        fm.pool.as_deref(),
1125    );
1126    let mut dg_pre = vec![0f32; rows * inter];
1127    let mut du_pre = vec![0f32; rows * inter];
1128    for row in 0..rows {
1129        for column in 0..inter {
1130            let at = row * inter + column;
1131            let da = dact[at] * gate[column];
1132            dg_pre[at] = da * up_pre[at] * ops::silu_bwd(gate_pre[at]);
1133            du_pre[at] = da * ops::silu(gate_pre[at]);
1134        }
1135    }
1136    let mut dg = vec![0f32; inter * hidden];
1137    let mut du = vec![0f32; inter * hidden];
1138    let mut dd = vec![0f32; hidden * inter];
1139    ops::gemm_dw(
1140        &dg_pre,
1141        &n2,
1142        &mut dg,
1143        rows,
1144        hidden,
1145        inter,
1146        fm.pool.as_deref(),
1147    );
1148    ops::gemm_dw(
1149        &du_pre,
1150        &n2,
1151        &mut du,
1152        rows,
1153        hidden,
1154        inter,
1155        fm.pool.as_deref(),
1156    );
1157    ops::gemm_dw(&dh, &act, &mut dd, rows, inter, hidden, fm.pool.as_deref());
1158    (
1159        score,
1160        Some((
1161            dg.into_iter().map(f64::from).collect(),
1162            du.into_iter().map(f64::from).collect(),
1163            dd.into_iter().map(f64::from).collect(),
1164        )),
1165    )
1166}
1167
1168/// Score a WRITTEN specialist through the replica's own math (f32
1169/// dequant of whatever the file carries) with the file's binary mask
1170/// held hard — the decomposition probe that tells "the requant at write
1171/// cost the quality" from "the runtime applies the mask differently".
1172/// Returns (bare, masked) held-PPL over the chunks.
1173pub fn replica_score_file_mask(
1174    model: &Arc<CmfModel>,
1175    chunks: &[Vec<u32>],
1176) -> Result<(f64, f64), String> {
1177    let o1_off = crate::nystrom::O1Cfg {
1178        layers: crate::nystrom::O1Layers::List(Vec::new()),
1179        m: 4,
1180        w: 8,
1181        sink: 1,
1182        rect: crate::nystrom::O1_DEFAULT_RECT,
1183    };
1184    let fm = FcdModel::from_cmf(model, &o1_off, false)?;
1185    let nl = fm.layers.len();
1186    let loops = fm.loops.max(1);
1187    let vn = nl * loops;
1188    let inter = fm.layers[0].inter;
1189    let ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
1190    // Binary mask → logits at ±50: σ crosses any τ exactly as the bit says.
1191    let task = &model.masks.default_task;
1192    let mask = model
1193        .masks
1194        .masks
1195        .iter()
1196        .find(|m| &m.name == task)
1197        .or_else(|| model.masks.masks.first());
1198    let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
1199    let masked_logits: Vec<Vec<f32>> = match mask {
1200        Some(m) => (0..vn)
1201            .map(|vl| {
1202                let row = m.ffn_masks.get(vl).map(|v| v.as_slice()).unwrap_or(&[]);
1203                (0..inter)
1204                    .map(|j| {
1205                        if (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0 {
1206                            50.0
1207                        } else {
1208                            -50.0
1209                        }
1210                    })
1211                    .collect()
1212            })
1213            .collect(),
1214        None => open.clone(),
1215    };
1216    let score = |logits: &[Vec<f32>]| -> f64 {
1217        let pass = Pass {
1218            fm: &fm,
1219            tau: 0.5,
1220            logits,
1221            hard: true,
1222            ffn: &ffn,
1223            focus_tokens: &[],
1224            focus_follow_tokens: &[],
1225        };
1226        held_ppl(&pass, chunks)
1227    };
1228    Ok((score(&open), score(&masked_logits)))
1229}
1230
1231/// The whole recipe. `log` receives progress lines.
1232pub fn skill_bake(
1233    model: &Arc<CmfModel>,
1234    chunks: &[Vec<u32>],
1235    held_n: usize,
1236    hy: &BakeHyper,
1237    mut log: impl FnMut(&str),
1238) -> Result<(BakeReport, BakeArtifacts), String> {
1239    let t0 = std::time::Instant::now();
1240    let o1_off = crate::nystrom::O1Cfg {
1241        layers: crate::nystrom::O1Layers::List(Vec::new()),
1242        m: 4,
1243        w: 8,
1244        sink: 1,
1245        rect: crate::nystrom::O1_DEFAULT_RECT,
1246    };
1247    let fm = FcdModel::from_cmf(model, &o1_off, false)?;
1248    let nl = fm.layers.len();
1249    let inter = fm.layers.iter().map(|l| l.inter).max().unwrap_or(0);
1250    let held: Vec<Vec<u32>> = chunks[..held_n.min(chunks.len())].to_vec();
1251    let calib: Vec<Vec<u32>> = chunks[held_n.min(chunks.len())..].to_vec();
1252    if calib.len() < 12 {
1253        return Err(format!(
1254            "skill bake: corpus too small ({} calib chunks)",
1255            calib.len()
1256        ));
1257    }
1258    // Analysis-only callers can request a native, batched focused score over
1259    // a validation corpus without training or writing a specialist.  The
1260    // scoring path already handles per-layer FFN widths; the optimizer and
1261    // defragment writer below still require one common width, so return the
1262    // identity checkpoint before that training-only constraint is applied.
1263    if hy.steps_a == 0 && hy.steps_b == 0 && hy.fcd_layers == 0 {
1264        let logits: Vec<Vec<f32>> = fm
1265            .layers
1266            .iter()
1267            .map(|layer| vec![100.0; layer.inter])
1268            .collect();
1269        let ffn = vec![None; nl];
1270        let pass = Pass {
1271            fm: &fm,
1272            tau: hy.tau,
1273            logits: &logits,
1274            hard: true,
1275            ffn: &ffn,
1276            focus_tokens: &hy.focus_tokens,
1277            focus_follow_tokens: &hy.focus_follow_tokens,
1278        };
1279        let score = held_score(&pass, &held);
1280        let keep: Vec<Vec<bool>> = fm
1281            .layers
1282            .iter()
1283            .map(|layer| vec![true; layer.inter])
1284            .collect();
1285        let loops = fm.loops.max(1);
1286        let keep_visits = (0..loops).flat_map(|_| keep.iter().cloned()).collect();
1287        let report = BakeReport {
1288            backbone: score.ppl,
1289            masked: score.ppl,
1290            overlaid: score.ppl,
1291            pruned_ratio: 0.0,
1292            kept_per_layer: keep.iter().map(Vec::len).collect(),
1293            backbone_accuracy: score.accuracy,
1294            masked_accuracy: score.accuracy,
1295            overlaid_accuracy: score.accuracy,
1296            backbone_balanced_accuracy: score.balanced_accuracy,
1297            masked_balanced_accuracy: score.balanced_accuracy,
1298            overlaid_balanced_accuracy: score.balanced_accuracy,
1299            selected_step: 0,
1300            sec: t0.elapsed().as_secs_f64(),
1301        };
1302        let arts = BakeArtifacts {
1303            keep,
1304            keep_visits,
1305            down: vec![Vec::new(); nl],
1306            gate_up: vec![None; nl],
1307            fcd_layers: Vec::new(),
1308            logits: logits.clone(),
1309            final_logits: logits,
1310            checkpoints: Vec::new(),
1311        };
1312        return Ok((report, arts));
1313    }
1314    if fm.layers.iter().any(|l| l.inter != inter) {
1315        return Err("skill bake: non-uniform FFN widths".into());
1316    }
1317    if hy.batch == 0 {
1318        return Err("skill bake: batch must be positive".into());
1319    }
1320    let fcd: Vec<usize> = (nl.saturating_sub(hy.fcd_layers)..nl).collect();
1321    let _rng = SplitMix64::new(hy.seed);
1322
1323    // Trainables. The gate starts as close to OPEN as the arithmetic
1324    // allows, because step zero must be the backbone and nothing else.
1325    //
1326    // It used to start at 2.0, and σ(2.0) = 0.881 — every FFN neuron
1327    // scaled to seven eighths before a single gradient. On an ordinary
1328    // stack that costs a few percent of perplexity and hides. On a
1329    // LOOPED Transformer it does not: Nanbeige 4.2 runs its 22 layers
1330    // twice, so each physical FFN is visited twice and the factor
1331    // compounds to 0.881² = 0.776 per layer over 44 visits. Measured:
1332    // baseline 4.187 → 278.4 at step 30 with 0% pruned. Nothing had been
1333    // pruned; the mask had simply turned the model down.
1334    //
1335    // So solve for the init instead of hardcoding it — but solve for the
1336    // right target. Pushing σ(m0)^loops to 0.999 starts at the backbone
1337    // and cannot move: the gradient carries σ'(m) = σ(1−σ), which at
1338    // σ = 0.9995 is 5e-4 against 0.105 at the old 2.0, and BOTH the data
1339    // term and the L1 term are scaled by it (see the update below). That
1340    // was tried: 60 steps, 0% pruned, hard-PPL equal to the baseline to
1341    // three digits. Identity that cannot learn is not an improvement.
1342    //
1343    // The quantity to preserve is the EFFECTIVE start — what the stack
1344    // actually multiplies by, once per visit compounded over the loop —
1345    // at the value the recipe was validated with on ordinary models:
1346    // σ(m0)^loops = σ(2.0) = 0.881. One loop reproduces the old constant
1347    // exactly, so nothing regresses; two loops open the per-visit gate to
1348    // 0.9385 so the compounded factor is again 0.881, with σ' = 0.058
1349    // rather than 0.0005.
1350    let loops = fm.loops.max(1);
1351    let m0 = mask_init_logit_for(loops, hy.mask_init);
1352    // One mask row per VIRTUAL layer: nl × loops. Unlooped: vn == nl.
1353    let vn = nl * loops;
1354    let mut logits: Vec<Vec<f32>> = vec![vec![m0; inter]; vn];
1355    let mut ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
1356
1357    // Baseline (no mask): even σ(m0) is not exactly 1, so measure with
1358    // gates forced open via hard mask over +∞… simplest: logits +50.
1359    let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
1360    let base_pass = Pass {
1361        fm: &fm,
1362        tau: hy.tau,
1363        logits: &open,
1364        hard: true,
1365        ffn: &ffn,
1366        focus_tokens: &hy.focus_tokens,
1367        focus_follow_tokens: &hy.focus_follow_tokens,
1368    };
1369    let backbone_score = held_score(&base_pass, &held);
1370    let backbone = backbone_score.ppl;
1371    log(&format!(
1372        "baseline (full): {backbone:.3}{}",
1373        backbone_score
1374            .accuracy
1375            .zip(backbone_score.balanced_accuracy)
1376            .map(|(a, b)| format!(" | acc {:.2}% bal {:.2}%", a * 100.0, b * 100.0))
1377            .unwrap_or_default()
1378    ));
1379
1380    // ── Phase A: mask training ──
1381    let mut adam_a = Adam::new(&vec![inter; vn], hy.lr_a);
1382    let mut l1 = hy.l1_init * hy.l1_mult;
1383    let l1_step_eff = hy.l1_step * hy.l1_mult;
1384    // best = (ppl, logits_snapshot, sparsity)
1385    let mut best: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
1386    let mut best_accuracy = f64::NEG_INFINITY;
1387    let mut best_balanced = f64::NEG_INFINITY;
1388    let mut best_step = 0usize;
1389    let mut checkpoints = Vec::new();
1390    // Track the highest-sparsity checkpoint as fallback.
1391    let mut max_sp: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
1392    let mut max_sp_step = 0usize;
1393    let mut prev_alive: Option<Vec<Vec<bool>>> = None;
1394    // In-process phase timers: this loop was estimated three different
1395    // ways and every estimate came out under a fifth of the measured
1396    // step time. Measure, then optimize the top line, not the guess.
1397    let mut acc_chunk = 0f64;
1398    let mut acc_adam = 0f64;
1399    // Phase A trains in strict f32: the mask SELECTS neurons by its
1400    // gradient, and f16 operand rounding on that signal — fine for every
1401    // forward and eval in this file — compounds over ~90 steps into
1402    // closing the wrong neurons (measured: hard-PPL 5.207 vs 4.293 at the
1403    // same 2.56% sparsity; the f32 run retraces the reference trajectory
1404    // to the third decimal). Evals inside the loop lift the restriction —
1405    // their tensor-core numbers match f32 at print precision.
1406    crate::gpu::bake_precision_strict(true);
1407    for step in 0..hy.steps_a {
1408        let t_step = std::time::Instant::now();
1409        let (batch_ids, batch) = calibration_batch(&calib, step, hy.batch)?;
1410        let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
1411        let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = vec![None; nl];
1412        let pass = Pass {
1413            fm: &fm,
1414            tau: hy.tau,
1415            logits: &logits,
1416            hard: false,
1417            ffn: &ffn,
1418            focus_tokens: &hy.focus_tokens,
1419            focus_follow_tokens: &hy.focus_follow_tokens,
1420        };
1421        let _ = pass.chunk_batch(&batch_ids, batch, Some((&mut dmask, &mut dffn)));
1422        // Fold σ'(m) into the mask grads + add the L1 term.
1423        let l1_per = l1 / (inter as f64 * nl as f64);
1424        for li in 0..vn {
1425            for j in 0..inter {
1426                let s = sigmoid(logits[li][j]) as f64;
1427                let sparse_grad = sparsity_grad(logits[li][j], hy.softplus_l1);
1428                dmask[li][j] = dmask[li][j] * s * (1.0 - s) + l1_per * sparse_grad;
1429            }
1430        }
1431        // One gradient per VISIT, one step per token. A Looped
1432        // Transformer visits each physical layer `loops` times and the
1433        // backward accumulates every visit into the same mask, so an
1434        // unnormalised step is `loops` times the one the recipe was
1435        // tuned with — the mask overshoots, neurons cross tau within the
1436        // first evaluation window, and each of them is then missing from
1437        // both passes. Dividing by the visit count makes a step mean the
1438        // same thing at any loop depth.
1439        // Per-visit rows: each mask logit receives exactly one visit's
1440        // gradient, so the step needs no visit normalisation here — that
1441        // scale now belongs to Phase B alone, where the FFN weights ARE
1442        // shared across visits.
1443        let t_chunk = t_step.elapsed().as_secs_f64();
1444        let mut params: Vec<&mut [f32]> = logits.iter_mut().map(|v| v.as_mut_slice()).collect();
1445        adam_a.step(&mut params, &dmask, 1.0);
1446        acc_chunk += t_chunk;
1447        acc_adam += t_step.elapsed().as_secs_f64() - t_chunk;
1448        if (step + 1) % hy.eval_every == 0 {
1449            l1 += l1_step_eff;
1450            let pass = Pass {
1451                fm: &fm,
1452                tau: hy.tau,
1453                logits: &logits,
1454                hard: true,
1455                ffn: &ffn,
1456                focus_tokens: &hy.focus_tokens,
1457                focus_follow_tokens: &hy.focus_follow_tokens,
1458            };
1459            crate::gpu::bake_precision_strict(false);
1460            let hs = held_score(&pass, &held);
1461            let hp = hs.ppl;
1462            crate::gpu::bake_precision_strict(true);
1463            // Name the neurons that crossed τ since the last eval. At
1464            // 0.01% pruned = ~24 neurons for a 135 held-PPL, WHICH 24 is
1465            // the whole diagnosis: it decides between "this model has no
1466            // noise neurons" and "a shared mask cannot spare a neuron
1467            // that only one visit of the loop needs".
1468            let cur: Vec<Vec<bool>> = logits
1469                .iter()
1470                .map(|l| l.iter().map(|&x| sigmoid(x) > hy.tau).collect())
1471                .collect();
1472            if let Some(prev) = &prev_alive {
1473                let died: Vec<String> = cur
1474                    .iter()
1475                    .zip(prev)
1476                    .enumerate()
1477                    .flat_map(|(li, (c, p))| {
1478                        c.iter()
1479                            .zip(p.iter())
1480                            .enumerate()
1481                            .filter(|&(_, (&cj, &pj))| pj && !cj)
1482                            .map(move |(j, _)| format!("L{li}:{j}"))
1483                    })
1484                    .collect();
1485                if !died.is_empty() {
1486                    log(&format!(
1487                        "    closed since last eval: {}: {}{}",
1488                        died.len(),
1489                        died.iter().take(32).cloned().collect::<Vec<_>>().join(" "),
1490                        if died.len() > 32 { " …" } else { "" }
1491                    ));
1492                }
1493            }
1494            let alive: usize = cur.iter().map(|l| l.iter().filter(|&&b| b).count()).sum();
1495            prev_alive = Some(cur);
1496            let sp = 1.0 - alive as f64 / (vn * inter) as f64;
1497            // Track highest-sparsity checkpoint.
1498            if sp > max_sp.2 {
1499                max_sp = (hp, Some(logits.clone()), sp);
1500                max_sp_step = step + 1;
1501            }
1502            checkpoints.push(BakeCheckpoint {
1503                step: step + 1,
1504                l1,
1505                ppl: hp,
1506                sparsity: sp,
1507                accuracy: hs.accuracy,
1508                balanced_accuracy: hs.balanced_accuracy,
1509            });
1510            // Best checkpoint selection: respect target_sparsity and any
1511            // caller-declared natural-distribution quality guards. The
1512            // guards are strict (`>`, not `>=`) so a majority baseline cannot
1513            // sneak through on an exactly tied checkpoint.
1514            let eligible_sparsity = hy.target_sparsity <= 0.0 || sp >= hy.target_sparsity;
1515            let eligible_accuracy = hy
1516                .checkpoint_min_accuracy
1517                .map_or(true, |min| hs.accuracy.is_some_and(|value| value > min));
1518            let eligible_balanced = hy.checkpoint_min_balanced_accuracy.map_or(true, |min| {
1519                hs.balanced_accuracy.is_some_and(|value| value > min)
1520            });
1521            let eligible = eligible_sparsity && eligible_accuracy && eligible_balanced;
1522            if eligible && hy.checkpoint_raw_priority && !hy.focus_tokens.is_empty() {
1523                let acc = hs.accuracy.unwrap_or(f64::NEG_INFINITY);
1524                let bal = hs.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1525                if acc > best_accuracy
1526                    || (acc == best_accuracy && bal > best_balanced)
1527                    || (acc == best_accuracy && bal == best_balanced && hp < best.0)
1528                {
1529                    best_balanced = bal;
1530                    best_accuracy = acc;
1531                    best = (hp, Some(logits.clone()), sp);
1532                    best_step = step + 1;
1533                }
1534            } else if eligible && hy.checkpoint_accuracy && !hy.focus_tokens.is_empty() {
1535                let acc = hs.accuracy.unwrap_or(f64::NEG_INFINITY);
1536                let bal = hs.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1537                if bal > best_balanced
1538                    || (bal == best_balanced && acc > best_accuracy)
1539                    || (bal == best_balanced && acc == best_accuracy && hp < best.0)
1540                {
1541                    best_balanced = bal;
1542                    best_accuracy = acc;
1543                    best = (hp, Some(logits.clone()), sp);
1544                    best_step = step + 1;
1545                }
1546            } else if eligible && hp < best.0 {
1547                best = (hp, Some(logits.clone()), sp);
1548                best_step = step + 1;
1549            }
1550            log(&format!(
1551                "  [A] step {}: L1={l1:.3} pruned={:.2}% hard-PPL={hp:.3}{} (bottom {}@{:.2}%) [fwd+bwd {:.1}s, adam {:.2}s per step]",
1552                step + 1,
1553                sp * 100.0,
1554                hs.accuracy
1555                    .zip(hs.balanced_accuracy)
1556                    .map(|(a, b)| format!(" acc={:.2}% bal={:.2}%", a * 100.0, b * 100.0))
1557                    .unwrap_or_default(),
1558                if best.0 == f64::MAX {
1559                    "—".to_string()
1560                } else {
1561                    format!("{:.3}", best.0)
1562                },
1563                best.2 * 100.0,
1564                acc_chunk / (step + 1) as f64,
1565                acc_adam / (step + 1) as f64
1566            ));
1567        }
1568    }
1569    // If target_sparsity was set but no checkpoint qualified, fall back
1570    // to the highest-sparsity checkpoint.
1571    // Phase A is over — phase B and every eval after run on the fast arms.
1572    crate::gpu::bake_precision_strict(false);
1573    if (hy.checkpoint_min_accuracy.is_some() || hy.checkpoint_min_balanced_accuracy.is_some())
1574        && best.1.is_none()
1575    {
1576        return Err(
1577            "skill bake: no Phase-A checkpoint met the configured focused accuracy guards".into(),
1578        );
1579    }
1580    if hy.target_sparsity > 0.0 && best.1.is_none() {
1581        log(&format!(
1582            "[A] target sparsity {:.0}% not reached; using max-sparsity checkpoint ({:.0}%)",
1583            hy.target_sparsity * 100.0,
1584            max_sp.2 * 100.0
1585        ));
1586        best = max_sp;
1587        best_step = max_sp_step;
1588    }
1589    // Phase totals, printed unconditionally — a 5-step measurement run
1590    // must report even though no eval fired.
1591    {
1592        use crate::fcd::prof;
1593        let (a, f, bw, g, gc) = (
1594            prof::take(&prof::ATTN_FWD),
1595            prof::take(&prof::FFN_FWD),
1596            prof::take(&prof::BWD),
1597            prof::take(&prof::GEMM),
1598            prof::GEMM_CALLS.swap(0, std::sync::atomic::Ordering::Relaxed),
1599        );
1600        log(&format!(
1601            "[prof] phase A over {} step(s): attn-fwd {a:.1}s | ffn-fwd {f:.1}s | bwd {bw:.1}s |              gemm total {g:.1}s in {gc} calls ({:.1} ms/call)",
1602            hy.steps_a,
1603            if gc > 0 { g * 1000.0 / gc as f64 } else { 0.0 }
1604        ));
1605        log(&format!("[prof] gemm shapes:\n{}", prof::shape_report(6)));
1606    }
1607    let final_logits = logits.clone();
1608    if let Some(b) = best.1.take() {
1609        logits = b;
1610    }
1611    let pass = Pass {
1612        fm: &fm,
1613        tau: hy.tau,
1614        logits: &logits,
1615        hard: true,
1616        ffn: &ffn,
1617        focus_tokens: &hy.focus_tokens,
1618        focus_follow_tokens: &hy.focus_follow_tokens,
1619    };
1620    // With no mask-training steps the hard mask is still exactly all-open,
1621    // so rescoring the same validation batch is pure duplicate work.  This
1622    // matters for focused 512-token records where one gate can take minutes
1623    // on a laptop, and keeps short FCD sweeps practical without changing a
1624    // single number.
1625    let masked_score = if hy.steps_a == 0 {
1626        backbone_score.clone()
1627    } else {
1628        held_score(&pass, &held)
1629    };
1630    let masked = masked_score.ppl;
1631    log(&format!(
1632        "[A] {:.0}s: masked-PPL {masked:.3}",
1633        t0.elapsed().as_secs_f64()
1634    ));
1635
1636    // ── Phase B: FCD of the last N layers' FFN (hard mask active) ──
1637    for &li in &fcd {
1638        let p = format!("model.layers.{li}.");
1639        ffn[li] = Some((
1640            crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.gate_proj.weight"))
1641                .map_err(|e| format!("phase-B gate: {e}"))?,
1642            crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.up_proj.weight"))
1643                .map_err(|e| format!("phase-B up: {e}"))?,
1644            crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.down_proj.weight"))
1645                .map_err(|e| format!("phase-B down: {e}"))?,
1646        ));
1647    }
1648    let sizes: Vec<usize> = fcd
1649        .iter()
1650        .flat_map(|&li| {
1651            let (g, u, d) = ffn[li].as_ref().expect("phase-B masters");
1652            [g.len(), u.len(), d.len()]
1653        })
1654        .collect();
1655    let mut adam_b = Adam::new(&sizes, hy.lr_b);
1656    // The mask-only model is a real checkpoint too. If every FCD eval is
1657    // worse, restore `None` overlays rather than accidentally writing the
1658    // final (rejected) training step while reporting the mask-only PPL.
1659    let mut best_b: (
1660        HeldScore,
1661        Option<Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>>>,
1662    ) = (masked_score.clone(), Some(vec![None; nl]));
1663    let cached_phase_b =
1664        hy.steps_b > 0 && fcd.len() == 1 && loops == 1 && !hy.focus_tokens.is_empty();
1665    if cached_phase_b {
1666        let last = fcd[0];
1667        let pass = Pass {
1668            fm: &fm,
1669            tau: hy.tau,
1670            logits: &logits,
1671            hard: true,
1672            ffn: &ffn,
1673            focus_tokens: &hy.focus_tokens,
1674            focus_follow_tokens: &hy.focus_follow_tokens,
1675        };
1676        log(&format!(
1677            "[B-cache] extracting native CMF boundaries: {} train + {} held records",
1678            calib.len(),
1679            held.len()
1680        ));
1681        let train_cache = build_focused_fcd_cache(&pass, &calib, 32)?;
1682        let held_cache = build_focused_fcd_cache(&pass, &held, 32)?;
1683        if train_cache.len() < 12 || held_cache.len() == 0 {
1684            return Err(format!(
1685                "focused final-FFN cache is too small: {} train, {} held answers",
1686                train_cache.len(),
1687                held_cache.len()
1688            ));
1689        }
1690        let gate = pass.gates(last);
1691        let held_indices: Vec<usize> = (0..held_cache.len()).collect();
1692        let (initial_cached, _) = {
1693            let (g, u, d) = ffn[last].as_ref().expect("cached FCD master");
1694            cached_fcd_run(
1695                &fm,
1696                &held_cache,
1697                &held_indices,
1698                (g, u, d),
1699                &gate,
1700                &hy.focus_tokens,
1701                false,
1702            )
1703        };
1704        log(&format!(
1705            "[B-cache] ready: {} train / {} held | parity PPL {:.3} vs full {:.3}{}",
1706            train_cache.len(),
1707            held_cache.len(),
1708            initial_cached.ppl,
1709            masked_score.ppl,
1710            initial_cached
1711                .accuracy
1712                .map(|value| format!(" | acc {:.2}%", value * 100.0))
1713                .unwrap_or_default()
1714        ));
1715        if (initial_cached.ppl - masked_score.ppl).abs() > 5e-3
1716            || initial_cached.accuracy != masked_score.accuracy
1717        {
1718            return Err(format!(
1719                "focused final-FFN cache parity failed: PPL {:.6} vs {:.6}, accuracy {:?} vs {:?}",
1720                initial_cached.ppl,
1721                masked_score.ppl,
1722                initial_cached.accuracy,
1723                masked_score.accuracy
1724            ));
1725        }
1726        for step in 0..hy.steps_b {
1727            let count = hy.fcd_batch.min(train_cache.len());
1728            let indices: Vec<usize> = (0..count)
1729                .map(|offset| (step * count + offset) % train_cache.len())
1730                .collect();
1731            let (_, gradients) = {
1732                let (g, u, d) = ffn[last].as_ref().expect("cached FCD master");
1733                cached_fcd_run(
1734                    &fm,
1735                    &train_cache,
1736                    &indices,
1737                    (g, u, d),
1738                    &gate,
1739                    &hy.focus_tokens,
1740                    true,
1741                )
1742            };
1743            let (dg, du, dd) = gradients.expect("cached FCD requested gradients");
1744            let lr_scale =
1745                0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
1746            let (g, u, d) = ffn[last].as_mut().expect("cached FCD master");
1747            let mut params = vec![g.as_mut_slice(), u.as_mut_slice(), d.as_mut_slice()];
1748            let grads = vec![dg, du, dd];
1749            adam_b.step(&mut params, &grads, lr_scale);
1750            if (step + 1) % hy.eval_every == 0 {
1751                let cur = {
1752                    let (g, u, d) = ffn[last].as_ref().expect("cached FCD master");
1753                    cached_fcd_run(
1754                        &fm,
1755                        &held_cache,
1756                        &held_indices,
1757                        (g, u, d),
1758                        &gate,
1759                        &hy.focus_tokens,
1760                        false,
1761                    )
1762                    .0
1763                };
1764                let cur_bal = cur.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1765                let best_bal = best_b.0.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1766                let cur_acc = cur.accuracy.unwrap_or(f64::NEG_INFINITY);
1767                let best_acc = best_b.0.accuracy.unwrap_or(f64::NEG_INFINITY);
1768                let better = if hy.checkpoint_accuracy {
1769                    cur_bal > best_bal
1770                        || (cur_bal == best_bal && cur_acc > best_acc)
1771                        || (cur_bal == best_bal && cur_acc == best_acc && cur.ppl < best_b.0.ppl)
1772                } else {
1773                    cur.ppl < best_b.0.ppl
1774                };
1775                if better {
1776                    best_b = (cur.clone(), Some(ffn.clone()));
1777                }
1778                log(&format!(
1779                    "  [B-cache] step {}: held-PPL {:.3} acc={:.2}% bal={:.2}% (best {:.3})",
1780                    step + 1,
1781                    cur.ppl,
1782                    cur_acc * 100.0,
1783                    cur_bal * 100.0,
1784                    best_b.0.ppl
1785                ));
1786            }
1787        }
1788    } else {
1789        for step in 0..hy.steps_b {
1790            let (batch_ids, batch) = calibration_batch(&calib, step, hy.batch)?;
1791            let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
1792            let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = (0..nl)
1793                .map(|li| {
1794                    ffn[li].as_ref().map(|(g, u, d)| {
1795                        (vec![0.0; g.len()], vec![0.0; u.len()], vec![0.0; d.len()])
1796                    })
1797                })
1798                .collect();
1799            let pass = Pass {
1800                fm: &fm,
1801                tau: hy.tau,
1802                logits: &logits,
1803                hard: true,
1804                ffn: &ffn,
1805                focus_tokens: &hy.focus_tokens,
1806                focus_follow_tokens: &hy.focus_follow_tokens,
1807            };
1808            let _ = pass.chunk_batch(&batch_ids, batch, Some((&mut dmask, &mut dffn)));
1809            // Cosine LR.
1810            let lr_scale =
1811                0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
1812            let first_fcd = fcd[0];
1813            let mut params: Vec<&mut [f32]> = Vec::new();
1814            let mut grads: Vec<Vec<f64>> = Vec::new();
1815            for (off, slot) in ffn[first_fcd..].iter_mut().enumerate() {
1816                let li = first_fcd + off;
1817                let Some((g, u, d)) = slot.as_mut() else {
1818                    continue;
1819                };
1820                let (dg, du, dd) = dffn[li].take().unwrap();
1821                params.push(g.as_mut_slice());
1822                grads.push(dg);
1823                params.push(u.as_mut_slice());
1824                grads.push(du);
1825                params.push(d.as_mut_slice());
1826                grads.push(dd);
1827            }
1828            // Same visit normalisation as Phase A: dffn accumulates every
1829            // visit of a physical layer, and an FFN update perturbs BOTH
1830            // passes of the loop, so per-step damage is `loops` times what
1831            // lr_b was tuned for on ordinary stacks.
1832            adam_b.step(&mut params, &grads, lr_scale * mask_step_scale(loops));
1833            if (step + 1) % hy.eval_every == 0 {
1834                let pass = Pass {
1835                    fm: &fm,
1836                    tau: hy.tau,
1837                    logits: &logits,
1838                    hard: true,
1839                    ffn: &ffn,
1840                    focus_tokens: &hy.focus_tokens,
1841                    focus_follow_tokens: &hy.focus_follow_tokens,
1842                };
1843                let cur = held_score(&pass, &held);
1844                let better = if hy.checkpoint_accuracy && !hy.focus_tokens.is_empty() {
1845                    let cur_bal = cur.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1846                    let best_bal = best_b.0.balanced_accuracy.unwrap_or(f64::NEG_INFINITY);
1847                    let cur_acc = cur.accuracy.unwrap_or(f64::NEG_INFINITY);
1848                    let best_acc = best_b.0.accuracy.unwrap_or(f64::NEG_INFINITY);
1849                    cur_bal > best_bal
1850                        || (cur_bal == best_bal && cur_acc > best_acc)
1851                        || (cur_bal == best_bal && cur_acc == best_acc && cur.ppl < best_b.0.ppl)
1852                } else {
1853                    cur.ppl < best_b.0.ppl
1854                };
1855                if better {
1856                    best_b = (cur.clone(), Some(ffn.clone()));
1857                }
1858                log(&format!(
1859                    "  [B] step {}: held-PPL {:.3}{} (best {:.3})",
1860                    step + 1,
1861                    cur.ppl,
1862                    cur.accuracy
1863                        .zip(cur.balanced_accuracy)
1864                        .map(|(a, b)| format!(" acc={:.2}% bal={:.2}%", a * 100.0, b * 100.0))
1865                        .unwrap_or_default(),
1866                    best_b.0.ppl
1867                ));
1868            }
1869        }
1870    }
1871    ffn = best_b.1.take().expect("phase-B always has a checkpoint");
1872    let overlaid_score = best_b.0;
1873    let overlaid = overlaid_score.ppl;
1874
1875    // ── Export artifacts ──
1876    // Per-visit keep flags are the mask that ships; the PHYSICAL keep is
1877    // their union, because a weight row can only be removed from disk if
1878    // no visit needs it.
1879    let keep_visits = keep_masks(&logits, hy.tau, hy.align, hy.uniform_inter);
1880    let keep: Vec<Vec<bool>> = (0..nl)
1881        .map(|li| {
1882            (0..inter)
1883                .map(|j| (0..loops).any(|v| keep_visits[v * nl + li][j]))
1884                .collect()
1885        })
1886        .collect();
1887    if hy.align > 1 || hy.uniform_inter {
1888        let raw: usize = logits
1889            .iter()
1890            .map(|l| l.iter().filter(|&&x| sigmoid(x) > hy.tau).count())
1891            .sum();
1892        // Compare like with like: raw σ-counts are over the VIRTUAL
1893        // rows, so the padded count must be too — the union rows are
1894        // fewer and the subtraction would underflow.
1895        let padded: usize = keep_visits
1896            .iter()
1897            .map(|a| a.iter().filter(|&&x| x).count())
1898            .sum::<usize>()
1899            .saturating_sub(raw);
1900        log(&format!(
1901            "align: +{padded} neurons resurrected (align {}, uniform {})",
1902            hy.align, hy.uniform_inter
1903        ));
1904    }
1905    let mut down_out = Vec::with_capacity(nl);
1906    let mut gate_up = Vec::with_capacity(nl);
1907    let mut kept_per_layer = Vec::with_capacity(nl);
1908    for li in 0..nl {
1909        let alive = &keep[li];
1910        kept_per_layer.push(alive.iter().filter(|&&a| a).count());
1911        let mut down = match &ffn[li] {
1912            Some((_, _, d)) => d.clone(),
1913            None => fm.mats(li).expect("layer mats").down.clone(),
1914        };
1915        let hsz = fm.hidden;
1916        for r in 0..hsz {
1917            for (c, &a) in alive.iter().enumerate() {
1918                if !a {
1919                    down[r * inter + c] = 0.0;
1920                }
1921            }
1922        }
1923        gate_up.push(ffn[li].as_ref().map(|(g, u, _)| (g.clone(), u.clone())));
1924        down_out.push(down);
1925    }
1926    let total: usize = keep_visits
1927        .iter()
1928        .map(|a| a.iter().filter(|&&x| x).count())
1929        .sum();
1930    let report = BakeReport {
1931        backbone,
1932        masked,
1933        overlaid,
1934        pruned_ratio: 1.0 - total as f64 / (vn * inter) as f64,
1935        kept_per_layer,
1936        backbone_accuracy: backbone_score.accuracy,
1937        masked_accuracy: masked_score.accuracy,
1938        overlaid_accuracy: overlaid_score.accuracy,
1939        backbone_balanced_accuracy: backbone_score.balanced_accuracy,
1940        masked_balanced_accuracy: masked_score.balanced_accuracy,
1941        overlaid_balanced_accuracy: overlaid_score.balanced_accuracy,
1942        selected_step: best_step,
1943        sec: t0.elapsed().as_secs_f64(),
1944    };
1945    let arts = BakeArtifacts {
1946        keep,
1947        keep_visits,
1948        down: down_out,
1949        gate_up,
1950        fcd_layers: fcd,
1951        logits: logits.clone(),
1952        final_logits,
1953        checkpoints,
1954    };
1955    Ok((report, arts))
1956}
1957
1958/// Hard-threshold keep masks from the trained logits, then resurrect
1959/// the highest-logit pruned neurons until each layer's kept count is a
1960/// multiple of `align` (rounding UP — the resurrected neurons are the
1961/// ones the mask ranked closest to the threshold, so this only moves
1962/// toward the full backbone). `uniform` additionally raises every layer
1963/// to the max layer's aligned count. A layer with 0 live neurons gets
1964/// `align.max(1)` — the defrag writer rejects empty layers.
1965fn keep_masks(logits: &[Vec<f32>], tau: f32, align: usize, uniform: bool) -> Vec<Vec<bool>> {
1966    let inter = logits[0].len();
1967    let round = |n: usize| -> usize {
1968        let n = n.max(1);
1969        if align <= 1 {
1970            n.min(inter)
1971        } else {
1972            (n.div_ceil(align) * align).min(inter)
1973        }
1974    };
1975    let mut want: Vec<usize> = logits
1976        .iter()
1977        .map(|l| round(l.iter().filter(|&&x| sigmoid(x) > tau).count()))
1978        .collect();
1979    if uniform {
1980        let k = want.iter().copied().max().unwrap_or(inter);
1981        want = vec![k; logits.len()];
1982    }
1983    logits
1984        .iter()
1985        .zip(&want)
1986        .map(|(l, &k)| {
1987            let mut idx: Vec<usize> = (0..inter).collect();
1988            idx.sort_unstable_by(|&a, &b| l[b].total_cmp(&l[a]));
1989            let mut alive = vec![false; inter];
1990            for &i in idx.iter().take(k) {
1991                alive[i] = true;
1992            }
1993            alive
1994        })
1995        .collect()
1996}
1997
1998#[cfg(test)]
1999mod tests {
2000    use super::*;
2001
2002    fn kept(masks: &[Vec<bool>]) -> Vec<usize> {
2003        masks
2004            .iter()
2005            .map(|m| m.iter().filter(|&&a| a).count())
2006            .collect()
2007    }
2008
2009    #[test]
2010    fn terminal_focus_ignores_label_names_inside_the_prompt() {
2011        // DOWN and UP occur in the instruction, but only the final UP is an
2012        // assistant answer because it is immediately followed by im_end.
2013        let down = 10;
2014        let up = 11;
2015        let im_end = 99;
2016        let ids = [1, down, 2, up, 3, up, im_end, 4];
2017        let focus = [down, up];
2018        let follow = [im_end];
2019        assert!(!is_scored_target(&ids, 1, ids.len(), &focus, &follow));
2020        assert!(!is_scored_target(&ids, 3, ids.len(), &focus, &follow));
2021        assert!(is_scored_target(&ids, 5, ids.len(), &focus, &follow));
2022    }
2023
2024    #[test]
2025    fn configurable_mask_init_preserves_effective_gate_across_loops() {
2026        for effective_logit in [2.0, 4.0] {
2027            let target = sigmoid(effective_logit);
2028            for loops in [1usize, 2, 4] {
2029                let per_visit = sigmoid(mask_init_logit_for(loops, effective_logit));
2030                assert!((per_visit.powi(loops as i32) - target).abs() < 2e-6);
2031            }
2032        }
2033    }
2034
2035    #[test]
2036    fn softplus_penalty_keeps_a_gradient_near_an_open_gate() {
2037        let gate_penalty = sparsity_grad(4.0, false);
2038        let softplus_penalty = sparsity_grad(4.0, true);
2039        assert!(softplus_penalty > gate_penalty * 50.0);
2040        assert!((softplus_penalty - sigmoid(4.0) as f64).abs() < 1e-7);
2041    }
2042
2043    #[test]
2044    fn balanced_accuracy_exposes_majority_class_collapse() {
2045        let stats = FocusStats {
2046            total: 100,
2047            correct: 90,
2048            class_total: vec![90, 10],
2049            class_correct: vec![90, 0],
2050        };
2051        assert_eq!(stats.accuracy(), Some(0.9));
2052        assert_eq!(stats.balanced_accuracy(), Some(0.5));
2053    }
2054
2055    #[test]
2056    fn grouped_focus_stats_merge_is_additive() {
2057        let mut all = FocusStats {
2058            total: 3,
2059            correct: 2,
2060            class_total: vec![2, 1],
2061            class_correct: vec![1, 1],
2062        };
2063        let second = FocusStats {
2064            total: 4,
2065            correct: 3,
2066            class_total: vec![1, 3],
2067            class_correct: vec![1, 2],
2068        };
2069        all.merge(&second);
2070        assert_eq!(all.total, 7);
2071        assert_eq!(all.correct, 5);
2072        assert_eq!(all.class_total, vec![3, 4]);
2073        assert_eq!(all.class_correct, vec![2, 3]);
2074        assert_eq!(all.accuracy(), Some(5.0 / 7.0));
2075        assert_eq!(all.balanced_accuracy(), Some((2.0 / 3.0 + 3.0 / 4.0) / 2.0));
2076    }
2077
2078    #[test]
2079    fn calibration_batches_keep_records_independent_and_wrap_deterministically() {
2080        let records = vec![vec![1, 2], vec![3, 4], vec![5, 6]];
2081        assert_eq!(
2082            calibration_batch(&records, 0, 2).unwrap(),
2083            (vec![1, 2, 3, 4], 2)
2084        );
2085        assert_eq!(
2086            calibration_batch(&records, 1, 2).unwrap(),
2087            (vec![5, 6, 1, 2], 2)
2088        );
2089        assert!(calibration_batch(&[vec![1, 2], vec![3]], 0, 2).is_err());
2090    }
2091
2092    /// align=32 rounds each layer UP by resurrecting the largest
2093    /// pruned logits; the originally-alive set stays alive.
2094    #[test]
2095    fn keep_masks_aligns_up_and_preserves_alive() {
2096        let inter = 96;
2097        // Layer 0: 40 alive (logits > 0 → σ > 0.5), the rest ramp
2098        // below threshold so resurrection order is deterministic.
2099        let l0: Vec<f32> = (0..inter)
2100            .map(|i| if i < 40 { 1.0 } else { -1.0 - i as f32 * 0.01 })
2101            .collect();
2102        // Layer 1: 64 alive — already aligned, must stay exactly 64.
2103        let l1: Vec<f32> = (0..inter)
2104            .map(|i| if i < 64 { 2.0 } else { -3.0 })
2105            .collect();
2106        let masks = keep_masks(&[l0.clone(), l1], 0.5, 32, false);
2107        assert_eq!(kept(&masks), vec![64, 64]);
2108        // The 40 originally-alive stay; resurrected are the top pruned
2109        // logits (indices 40..64 — the least-negative of the ramp).
2110        for i in 0..64 {
2111            assert!(masks[0][i], "neuron {i} should be kept");
2112        }
2113        for i in 64..inter {
2114            assert!(!masks[0][i], "neuron {i} should stay pruned");
2115        }
2116    }
2117
2118    /// uniform=true raises every layer to the max aligned count.
2119    #[test]
2120    fn keep_masks_uniform_takes_max() {
2121        let inter = 96;
2122        let l0: Vec<f32> = (0..inter)
2123            .map(|i| if i < 10 { 1.0 } else { -2.0 })
2124            .collect();
2125        let l1: Vec<f32> = (0..inter)
2126            .map(|i| if i < 70 { 1.0 } else { -2.0 })
2127            .collect();
2128        let masks = keep_masks(&[l0, l1], 0.5, 32, true);
2129        assert_eq!(kept(&masks), vec![96, 96]);
2130    }
2131
2132    /// align capped at inter; align=1 (off) keeps the raw threshold
2133    /// count; an all-pruned layer still keeps at least one neuron.
2134    #[test]
2135    fn keep_masks_edges() {
2136        let inter = 48;
2137        let l: Vec<f32> = (0..inter)
2138            .map(|i| if i < 47 { 1.0 } else { -2.0 })
2139            .collect();
2140        let masks = keep_masks(&[l.clone()], 0.5, 32, false);
2141        assert_eq!(kept(&masks), vec![48]); // 47 → 64 capped to 48
2142        let masks = keep_masks(&[l], 0.5, 1, false);
2143        assert_eq!(kept(&masks), vec![47]);
2144        let dead: Vec<f32> = vec![-5.0; inter];
2145        let masks = keep_masks(&[dead], 0.5, 32, false);
2146        assert_eq!(kept(&masks), vec![32]); // max(1) → rounded to 32
2147    }
2148}