1use crate::fcd::{FcdModel, LnFfn};
22use crate::fcd_ops as ops;
23use crate::sampler::SplitMix64;
24use cortiq_core::CmfModel;
25use std::sync::Arc;
26
27#[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 pub batch: usize,
42 pub fcd_batch: usize,
46 pub seed: u64,
47 pub target_sparsity: f64,
51 pub l1_mult: f64,
54 pub mask_init: f32,
59 pub softplus_l1: bool,
63 pub checkpoint_accuracy: bool,
67 pub checkpoint_min_accuracy: Option<f64>,
72 pub checkpoint_min_balanced_accuracy: Option<f64>,
76 pub checkpoint_raw_priority: bool,
80 pub align: usize,
84 pub uniform_inter: bool,
87 pub focus_tokens: Vec<u32>,
93 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
131pub struct BakeReport {
133 pub backbone: f64,
135 pub masked: f64,
137 pub overlaid: f64,
139 pub pruned_ratio: f64,
140 pub kept_per_layer: Vec<usize>,
141 pub backbone_accuracy: Option<f64>,
143 pub masked_accuracy: Option<f64>,
144 pub overlaid_accuracy: Option<f64>,
145 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
163pub struct BakeArtifacts {
165 pub keep: Vec<Vec<bool>>,
168 pub keep_visits: Vec<Vec<bool>>,
171 pub down: Vec<Vec<f32>>,
174 pub gate_up: Vec<Option<(Vec<f32>, Vec<f32>)>>,
176 pub fcd_layers: Vec<usize>,
178 pub logits: Vec<Vec<f32>>,
183 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
194struct 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 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
236pub 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
264pub 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
297struct Pass<'a> {
300 fm: &'a FcdModel,
301 tau: f32,
302 logits: &'a [Vec<f32>],
304 hard: bool,
305 ffn: &'a [Option<(Vec<f32>, Vec<f32>, Vec<f32>)>],
307 focus_tokens: &'a [u32],
309 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#[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 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 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 #[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 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 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 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 let mut lnorms: Vec<Option<(Vec<f32>, Vec<f32>)>> = vec![None; vn];
579 for vl in 0..vn {
580 let li = vl % nl;
581 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 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 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]; 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 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 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 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 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 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 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 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 {
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 for (j, d) in dm.iter_mut().enumerate() {
816 let _ = j;
817 let _ = d;
818 }
819 }
820 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 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 let mut dh1 = dh.clone(); 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
891fn 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 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
976fn 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
1168pub 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 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
1231pub 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 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 let loops = fm.loops.max(1);
1351 let m0 = mask_init_logit_for(loops, hy.mask_init);
1352 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 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 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 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 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 let mut acc_chunk = 0f64;
1398 let mut acc_adam = 0f64;
1399 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 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 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 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 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 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 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 {
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 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 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 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 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 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 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 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
1958fn 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 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 #[test]
2095 fn keep_masks_aligns_up_and_preserves_alive() {
2096 let inter = 96;
2097 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 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 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 #[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 #[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]); 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]); }
2148}