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 seed: u64,
40 pub target_sparsity: f64,
44 pub l1_mult: f64,
47 pub align: usize,
51 pub uniform_inter: bool,
54}
55
56impl Default for BakeHyper {
57 fn default() -> Self {
58 Self {
59 steps_a: 240,
60 steps_b: 120,
61 l1_init: 0.01,
62 l1_step: 0.005,
63 eval_every: 30,
64 lr_a: 0.1,
65 lr_b: 1e-5,
66 tau: 0.5,
67 fcd_layers: 4,
68 seed: 0,
69 target_sparsity: 0.0,
70 l1_mult: 1.0,
71 align: 32,
72 uniform_inter: false,
73 }
74 }
75}
76
77pub struct BakeReport {
79 pub backbone: f64,
81 pub masked: f64,
83 pub overlaid: f64,
85 pub pruned_ratio: f64,
86 pub kept_per_layer: Vec<usize>,
87 pub sec: f64,
88}
89
90pub struct BakeArtifacts {
92 pub keep: Vec<Vec<bool>>,
95 pub keep_visits: Vec<Vec<bool>>,
98 pub down: Vec<Vec<f32>>,
101 pub gate_up: Vec<Option<(Vec<f32>, Vec<f32>)>>,
103 pub fcd_layers: Vec<usize>,
105}
106
107const CLIP: f64 = 1.0;
108const B1: f64 = 0.9;
109const B2: f64 = 0.999;
110const EPS: f64 = 1e-8;
111
112struct Adam {
114 m: Vec<Vec<f64>>,
115 v: Vec<Vec<f64>>,
116 t: i32,
117 lr: f64,
118}
119
120impl Adam {
121 fn new(sizes: &[usize], lr: f64) -> Self {
122 Self {
123 m: sizes.iter().map(|&n| vec![0.0; n]).collect(),
124 v: sizes.iter().map(|&n| vec![0.0; n]).collect(),
125 t: 0,
126 lr,
127 }
128 }
129
130 fn step(&mut self, params: &mut [&mut [f32]], grads: &[Vec<f64>], lr_scale: f64) {
132 let gn: f64 = grads
133 .iter()
134 .flat_map(|g| g.iter().map(|x| x * x))
135 .sum::<f64>()
136 .sqrt();
137 let clip = if gn > CLIP { CLIP / gn } else { 1.0 };
138 self.t += 1;
139 let (bc1, bc2) = (1.0 - B1.powi(self.t), 1.0 - B2.powi(self.t));
140 for (pi, p) in params.iter_mut().enumerate() {
141 for j in 0..p.len() {
142 let g = grads[pi][j] * clip;
143 let m = &mut self.m[pi][j];
144 let v = &mut self.v[pi][j];
145 *m = B1 * *m + (1.0 - B1) * g;
146 *v = B2 * *v + (1.0 - B2) * g * g;
147 let upd = (*m / bc1) / ((*v / bc2).sqrt() + EPS);
148 p[j] -= (self.lr * lr_scale * upd) as f32;
149 }
150 }
151 }
152}
153
154pub fn mask_init_logit(loops: usize) -> f32 {
173 let base = 1.0f32 / (1.0 + (-2.0f32).exp());
174 let per_visit = base.powf(1.0 / loops.max(1) as f32);
175 (per_visit / (1.0 - per_visit)).ln()
176}
177
178pub fn mask_step_scale(loops: usize) -> f64 {
184 1.0 / loops.max(1) as f64
185}
186
187fn sigmoid(x: f32) -> f32 {
188 1.0 / (1.0 + (-x).exp())
189}
190
191struct Pass<'a> {
194 fm: &'a FcdModel,
195 tau: f32,
196 logits: &'a [Vec<f32>],
198 hard: bool,
199 ffn: &'a [Option<(Vec<f32>, Vec<f32>, Vec<f32>)>],
201}
202
203impl Pass<'_> {
204 fn gates(&self, li: usize) -> Vec<f32> {
205 self.logits[li]
206 .iter()
207 .map(|&l| {
208 let s = sigmoid(l);
209 if self.hard {
210 if s > self.tau { 1.0 } else { 0.0 }
211 } else {
212 s
213 }
214 })
215 .collect()
216 }
217
218 fn wts<'b>(&'b self, li: usize) -> LnFfn<'b> {
219 let l = &self.fm.layers[li];
220 match &self.ffn[li] {
221 Some((g, u, d)) => LnFfn {
222 iln: &l.iln,
223 pln: &l.pln,
224 gate: g,
225 up: u,
226 down: d,
227 gu: None,
229 },
230 None => LnFfn {
231 iln: &l.iln,
232 pln: &l.pln,
233 gate: &l.gate,
234 up: &l.up,
235 down: &l.down,
236 gu: Some(&l.gu),
237 },
238 }
239 }
240
241 #[allow(clippy::too_many_arguments)]
245 fn chunk(
246 &self,
247 ids: &[u32],
248 grad: Option<(
249 &mut [Vec<f64>],
250 &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
251 )>,
252 ) -> (f64, usize) {
253 self.chunk_batch(ids, 1, grad)
254 }
255
256 fn chunk_batch(
264 &self,
265 ids: &[u32],
266 b: usize,
267 grad: Option<(
268 &mut [Vec<f64>],
269 &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
270 )>,
271 ) -> (f64, usize) {
272 let fm = self.fm;
273 let hsz = fm.hidden;
274 debug_assert!(ids.len() % b.max(1) == 0, "ragged batch");
275 debug_assert!(grad.is_none() || b == 1, "grads are per-chunk");
278 let t = ids.len() / b.max(1);
279 let n = b * t;
280 let nl = fm.layers.len();
281 let mut h = vec![0f32; n * hsz];
283 for (r, &id) in ids.iter().enumerate() {
284 h[r * hsz..(r + 1) * hsz]
285 .copy_from_slice(&fm.embed[id as usize * hsz..(id as usize + 1) * hsz]);
286 }
287 let loops = fm.loops.max(1);
295 let vn = nl * loops;
296 let mut h_ins = Vec::with_capacity(vn);
297 let mut acts = Vec::with_capacity(vn);
298 let mut masks = Vec::with_capacity(vn);
299 let mut lnorms: Vec<Option<(Vec<f32>, Vec<f32>)>> = vec![None; vn];
301 for vl in 0..vn {
302 let li = vl % nl;
303 let g = self.gates(vl);
308 let wts = self.wts(li);
309 let want = grad.is_some();
310 let (h2, a) = fm.layer_forward_scaled(li, &h, b, t, &wts, false, want, Some(&g));
311 h_ins.push(if want { h } else { Vec::new() });
312 acts.push(a);
313 masks.push(g);
314 h = h2;
315 if fm.loop_norm && li + 1 == nl && vl + 1 < vn {
318 let mut hn = vec![0f32; n * hsz];
319 let mut inv = vec![0f32; n];
320 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
321 if want {
322 lnorms[vl] = Some((h, inv));
323 }
324 h = hn;
325 }
326 }
327 let mut hn = vec![0f32; n * hsz];
329 let mut inv = vec![0f32; n];
330 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
331 let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
332 let vocab = lm.len() / hsz;
333 let pool = fm.pool.as_deref();
334 let mut nll = 0f64;
335 let mut dh_n = vec![0f32; n * hsz]; const POS_CHUNK: usize = 64;
340 let scored = b * (t - 1);
341 for bi in 0..b {
342 let base = bi * t;
343 let mut p0 = 0usize;
344 while p0 < t - 1 {
345 let pc = POS_CHUNK.min(t - 1 - p0);
346 let mut logits = vec![0f32; pc * vocab];
347 ops::gemm_nt(
348 &hn[(base + p0) * hsz..(base + p0 + pc) * hsz],
349 lm,
350 &mut logits,
351 pc,
352 hsz,
353 vocab,
354 pool,
355 );
356 for r in 0..pc {
357 let target = ids[base + p0 + r + 1] as usize;
358 let row = &mut logits[r * vocab..(r + 1) * vocab];
359 let mx = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
360 let mut sum = 0f64;
361 for v in row.iter() {
362 sum += ((*v as f64) - mx).exp();
363 }
364 nll += mx + sum.ln() - row[target] as f64;
365 if grad.is_some() {
366 let inv_n = 1.0 / scored as f64;
368 for v in row.iter_mut() {
369 *v = ((((*v as f64) - mx).exp() / sum) * inv_n) as f32;
370 }
371 row[target] -= inv_n as f32;
372 }
373 }
374 if grad.is_some() {
375 ops::gemm_dx(
376 &logits,
377 lm,
378 &mut dh_n[(base + p0) * hsz..(base + p0 + pc) * hsz],
379 pc,
380 hsz,
381 vocab,
382 pool,
383 );
384 }
385 p0 += pc;
386 }
387 }
388 let Some((dmask, dffn)) = grad else {
389 return (nll, scored);
390 };
391 let t_bwd = std::time::Instant::now();
393 let mut dh = vec![0f32; n * hsz];
394 ops::rmsnorm_bwd(&h, &fm.final_norm, &inv, &dh_n, fm.gemma, &mut dh, None);
395 for vl in (0..vn).rev() {
396 let li = vl % nl;
397 if let Some((hb, inv)) = lnorms[vl].as_ref() {
399 let mut dprev = vec![0f32; n * hsz];
400 ops::rmsnorm_bwd(hb, &fm.final_norm, inv, &dh, fm.gemma, &mut dprev, None);
401 dh = dprev;
402 }
403 let a = acts[vl].as_ref().expect("acts saved in grad mode");
404 let g = &masks[vl];
405 let inter = fm.layers[li].inter;
406 let wts = self.wts(li);
407 let mut dact2 = vec![0f32; t * inter];
409 ops::gemm_dx(&dh, wts.down, &mut dact2, t, inter, hsz, fm.pool.as_deref());
410 if let Some((_, _, dd)) = dffn[li].as_mut() {
411 let mut act2 = a.act.clone();
413 for r in 0..t {
414 for (x, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
415 *x *= gv;
416 }
417 }
418 let mut dw = vec![0f32; hsz * inter];
419 ops::gemm_dw(&dh, &act2, &mut dw, t, inter, hsz, fm.pool.as_deref());
420 for (o, &x) in dd.iter_mut().zip(&dw) {
421 *o += x as f64;
422 }
423 }
424 {
428 let dm = &mut dmask[vl];
429 for r in 0..t {
430 let da = &dact2[r * inter..(r + 1) * inter];
431 let aa = &a.act[r * inter..(r + 1) * inter];
432 for j in 0..inter {
433 dm[j] += da[j] as f64 * aa[j] as f64;
434 }
435 }
436 for (j, d) in dm.iter_mut().enumerate() {
438 let _ = j;
439 let _ = d;
440 }
441 }
442 let mut dg_pre = vec![0f32; t * inter];
444 let mut du_pre = vec![0f32; t * inter];
445 for r in 0..t {
446 for j in 0..inter {
447 let i = r * inter + j;
448 let da = dact2[i] * g[j];
449 let sg = ops::silu(a.gpre[i]);
450 dg_pre[i] = da * a.upre[i] * ops::silu_bwd(a.gpre[i]);
451 du_pre[i] = da * sg;
452 }
453 }
454 let mut dn2 = vec![0f32; t * hsz];
457 if let Some(gu) = wts.gu {
458 let mut dgu = vec![0f32; t * 2 * inter];
459 for r in 0..t {
460 let row = &mut dgu[r * 2 * inter..(r + 1) * 2 * inter];
461 row[..inter].copy_from_slice(&dg_pre[r * inter..(r + 1) * inter]);
462 row[inter..].copy_from_slice(&du_pre[r * inter..(r + 1) * inter]);
463 }
464 ops::gemm_dx(&dgu, gu, &mut dn2, t, hsz, 2 * inter, fm.pool.as_deref());
465 } else {
466 ops::gemm_dx(
467 &dg_pre,
468 wts.gate,
469 &mut dn2,
470 t,
471 hsz,
472 inter,
473 fm.pool.as_deref(),
474 );
475 let mut dn2b = vec![0f32; t * hsz];
476 ops::gemm_dx(
477 &du_pre,
478 wts.up,
479 &mut dn2b,
480 t,
481 hsz,
482 inter,
483 fm.pool.as_deref(),
484 );
485 for (x, &y) in dn2.iter_mut().zip(&dn2b) {
486 *x += y;
487 }
488 }
489 if let Some((dgw, duw, _)) = dffn[li].as_mut() {
490 let mut dw = vec![0f32; inter * hsz];
491 ops::gemm_dw(&dg_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
492 for (o, &x) in dgw.iter_mut().zip(&dw) {
493 *o += x as f64;
494 }
495 dw.fill(0.0);
496 ops::gemm_dw(&du_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
497 for (o, &x) in duw.iter_mut().zip(&dw) {
498 *o += x as f64;
499 }
500 }
501 let mut dh1 = dh.clone(); ops::rmsnorm_bwd(&a.h1, wts.pln, &a.inv2, &dn2, fm.gemma, &mut dh1, None);
505 dh = dh1;
506 let _ = &h_ins[vl];
507 }
508 crate::fcd::prof::add(&crate::fcd::prof::BWD, t_bwd);
509 (nll, scored)
510 }
511}
512
513fn held_ppl(pass: &Pass, held: &[Vec<u32>]) -> f64 {
515 if held.is_empty() {
519 return f64::NAN;
520 }
521 let t = held[0].len();
522 if held.iter().all(|c| c.len() == t) {
523 let flat: Vec<u32> = held.iter().flatten().copied().collect();
524 let (l, k) = pass.chunk_batch(&flat, held.len(), None);
525 return (l / k.max(1) as f64).exp();
526 }
527 let mut nll = 0f64;
528 let mut n = 0usize;
529 for c in held {
530 let (l, k) = pass.chunk(c, None);
531 nll += l;
532 n += k;
533 }
534 (nll / n.max(1) as f64).exp()
535}
536
537pub fn skill_bake(
539 model: &Arc<CmfModel>,
540 chunks: &[Vec<u32>],
541 held_n: usize,
542 hy: &BakeHyper,
543 mut log: impl FnMut(&str),
544) -> Result<(BakeReport, BakeArtifacts), String> {
545 let t0 = std::time::Instant::now();
546 let o1_off = crate::nystrom::O1Cfg {
547 layers: crate::nystrom::O1Layers::List(Vec::new()),
548 m: 4,
549 w: 8,
550 sink: 1,
551 rect: crate::nystrom::O1_DEFAULT_RECT,
552 };
553 let fm = FcdModel::from_cmf(model, &o1_off)?;
554 let nl = fm.layers.len();
555 let inter = fm.layers.iter().map(|l| l.inter).max().unwrap_or(0);
556 if fm.layers.iter().any(|l| l.inter != inter) {
557 return Err("skill bake: non-uniform FFN widths".into());
558 }
559 let held: Vec<Vec<u32>> = chunks[..held_n.min(chunks.len())].to_vec();
560 let calib: Vec<Vec<u32>> = chunks[held_n.min(chunks.len())..].to_vec();
561 if calib.len() < 12 {
562 return Err(format!(
563 "skill bake: corpus too small ({} calib chunks)",
564 calib.len()
565 ));
566 }
567 let fcd: Vec<usize> = (nl.saturating_sub(hy.fcd_layers)..nl).collect();
568 let _rng = SplitMix64::new(hy.seed);
569
570 let loops = fm.loops.max(1);
598 let m0 = mask_init_logit(loops);
599 let vn = nl * loops;
601 let mut logits: Vec<Vec<f32>> = vec![vec![m0; inter]; vn];
602 let mut ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
603
604 let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
607 let base_pass = Pass {
608 fm: &fm,
609 tau: hy.tau,
610 logits: &open,
611 hard: true,
612 ffn: &ffn,
613 };
614 let backbone = held_ppl(&base_pass, &held);
615 log(&format!("baseline (full): {backbone:.3}"));
616
617 let mut adam_a = Adam::new(&vec![inter; vn], hy.lr_a);
619 let mut l1 = hy.l1_init * hy.l1_mult;
620 let l1_step_eff = hy.l1_step * hy.l1_mult;
621 let mut best: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
623 let mut max_sp: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
625 let mut prev_alive: Option<Vec<Vec<bool>>> = None;
626 let mut acc_chunk = 0f64;
630 let mut acc_adam = 0f64;
631 for step in 0..hy.steps_a {
632 let t_step = std::time::Instant::now();
633 let chunk = &calib[step % calib.len()];
634 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
635 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = vec![None; nl];
636 let pass = Pass {
637 fm: &fm,
638 tau: hy.tau,
639 logits: &logits,
640 hard: false,
641 ffn: &ffn,
642 };
643 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
644 let l1_per = l1 / (inter as f64 * nl as f64);
646 for li in 0..vn {
647 for j in 0..inter {
648 let s = sigmoid(logits[li][j]) as f64;
649 dmask[li][j] = dmask[li][j] * s * (1.0 - s) + l1_per * s * (1.0 - s);
650 }
651 }
652 let t_chunk = t_step.elapsed().as_secs_f64();
665 let mut params: Vec<&mut [f32]> = logits.iter_mut().map(|v| v.as_mut_slice()).collect();
666 adam_a.step(&mut params, &dmask, 1.0);
667 acc_chunk += t_chunk;
668 acc_adam += t_step.elapsed().as_secs_f64() - t_chunk;
669 if (step + 1) % hy.eval_every == 0 {
670 l1 += l1_step_eff;
671 let pass = Pass {
672 fm: &fm,
673 tau: hy.tau,
674 logits: &logits,
675 hard: true,
676 ffn: &ffn,
677 };
678 let hp = held_ppl(&pass, &held);
679 let cur: Vec<Vec<bool>> = logits
685 .iter()
686 .map(|l| l.iter().map(|&x| sigmoid(x) > hy.tau).collect())
687 .collect();
688 if let Some(prev) = &prev_alive {
689 let died: Vec<String> = cur
690 .iter()
691 .zip(prev)
692 .enumerate()
693 .flat_map(|(li, (c, p))| {
694 c.iter()
695 .zip(p.iter())
696 .enumerate()
697 .filter(|&(_, (&cj, &pj))| pj && !cj)
698 .map(move |(j, _)| format!("L{li}:{j}"))
699 })
700 .collect();
701 if !died.is_empty() {
702 log(&format!(
703 " closed since last eval: {}: {}{}",
704 died.len(),
705 died.iter().take(32).cloned().collect::<Vec<_>>().join(" "),
706 if died.len() > 32 { " …" } else { "" }
707 ));
708 }
709 }
710 let alive: usize = cur.iter().map(|l| l.iter().filter(|&&b| b).count()).sum();
711 prev_alive = Some(cur);
712 let sp = 1.0 - alive as f64 / (vn * inter) as f64;
713 if sp > max_sp.2 {
715 max_sp = (hp, Some(logits.clone()), sp);
716 }
717 if hy.target_sparsity > 0.0 {
719 if sp >= hy.target_sparsity && hp < best.0 {
720 best = (hp, Some(logits.clone()), sp);
721 }
722 } else if hp < best.0 {
723 best = (hp, Some(logits.clone()), sp);
724 }
725 log(&format!(
726 " [A] step {}: L1={l1:.3} pruned={:.2}% hard-PPL={hp:.3} (bottom {}@{:.2}%) [fwd+bwd {:.1}s, adam {:.2}s per step]",
727 step + 1,
728 sp * 100.0,
729 if best.0 == f64::MAX {
730 "—".to_string()
731 } else {
732 format!("{:.3}", best.0)
733 },
734 best.2 * 100.0,
735 acc_chunk / (step + 1) as f64,
736 acc_adam / (step + 1) as f64
737 ));
738 }
739 }
740 if hy.target_sparsity > 0.0 && best.1.is_none() {
743 log(&format!(
744 "[A] target sparsity {:.0}% not reached; using max-sparsity checkpoint ({:.0}%)",
745 hy.target_sparsity * 100.0,
746 max_sp.2 * 100.0
747 ));
748 best = max_sp;
749 }
750 {
753 use crate::fcd::prof;
754 let (a, f, bw, g, gc) = (
755 prof::take(&prof::ATTN_FWD),
756 prof::take(&prof::FFN_FWD),
757 prof::take(&prof::BWD),
758 prof::take(&prof::GEMM),
759 prof::GEMM_CALLS.swap(0, std::sync::atomic::Ordering::Relaxed),
760 );
761 log(&format!(
762 "[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)",
763 hy.steps_a,
764 if gc > 0 { g * 1000.0 / gc as f64 } else { 0.0 }
765 ));
766 log(&format!("[prof] gemm shapes:\n{}", prof::shape_report(6)));
767 }
768 if let Some(b) = best.1.take() {
769 logits = b;
770 }
771 let pass = Pass {
772 fm: &fm,
773 tau: hy.tau,
774 logits: &logits,
775 hard: true,
776 ffn: &ffn,
777 };
778 let masked = held_ppl(&pass, &held);
779 log(&format!(
780 "[A] {:.0}s: masked-PPL {masked:.3}",
781 t0.elapsed().as_secs_f64()
782 ));
783
784 for &li in &fcd {
786 let l = &fm.layers[li];
787 ffn[li] = Some((l.gate.clone(), l.up.clone(), l.down.clone()));
788 }
789 let sizes: Vec<usize> = fcd
790 .iter()
791 .flat_map(|&li| {
792 let l = &fm.layers[li];
793 [l.gate.len(), l.up.len(), l.down.len()]
794 })
795 .collect();
796 let mut adam_b = Adam::new(&sizes, hy.lr_b);
797 let mut best_b: (f64, Option<Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>>>) = (masked, None);
798 for step in 0..hy.steps_b {
799 let chunk = &calib[step % calib.len()];
800 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
801 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = (0..nl)
802 .map(|li| {
803 ffn[li]
804 .as_ref()
805 .map(|(g, u, d)| (vec![0.0; g.len()], vec![0.0; u.len()], vec![0.0; d.len()]))
806 })
807 .collect();
808 let pass = Pass {
809 fm: &fm,
810 tau: hy.tau,
811 logits: &logits,
812 hard: true,
813 ffn: &ffn,
814 };
815 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
816 let lr_scale = 0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
818 let first_fcd = fcd[0];
819 let mut params: Vec<&mut [f32]> = Vec::new();
820 let mut grads: Vec<Vec<f64>> = Vec::new();
821 for (off, slot) in ffn[first_fcd..].iter_mut().enumerate() {
822 let li = first_fcd + off;
823 let Some((g, u, d)) = slot.as_mut() else {
824 continue;
825 };
826 let (dg, du, dd) = dffn[li].take().unwrap();
827 params.push(g.as_mut_slice());
828 grads.push(dg);
829 params.push(u.as_mut_slice());
830 grads.push(du);
831 params.push(d.as_mut_slice());
832 grads.push(dd);
833 }
834 adam_b.step(&mut params, &grads, lr_scale * mask_step_scale(loops));
839 if (step + 1) % hy.eval_every == 0 {
840 let pass = Pass {
841 fm: &fm,
842 tau: hy.tau,
843 logits: &logits,
844 hard: true,
845 ffn: &ffn,
846 };
847 let cur = held_ppl(&pass, &held);
848 if cur < best_b.0 {
849 best_b = (cur, Some(ffn.clone()));
850 }
851 log(&format!(
852 " [B] step {}: held-PPL {cur:.3} (best {:.3})",
853 step + 1,
854 best_b.0
855 ));
856 }
857 }
858 if let Some(b) = best_b.1.take() {
859 ffn = b;
860 }
861 let overlaid = best_b.0;
862
863 let keep_visits = keep_masks(&logits, hy.tau, hy.align, hy.uniform_inter);
868 let keep: Vec<Vec<bool>> = (0..nl)
869 .map(|li| {
870 (0..inter)
871 .map(|j| (0..loops).any(|v| keep_visits[v * nl + li][j]))
872 .collect()
873 })
874 .collect();
875 if hy.align > 1 || hy.uniform_inter {
876 let raw: usize = logits
877 .iter()
878 .map(|l| l.iter().filter(|&&x| sigmoid(x) > hy.tau).count())
879 .sum();
880 let padded: usize = keep_visits
884 .iter()
885 .map(|a| a.iter().filter(|&&x| x).count())
886 .sum::<usize>()
887 .saturating_sub(raw);
888 log(&format!(
889 "align: +{padded} neurons resurrected (align {}, uniform {})",
890 hy.align, hy.uniform_inter
891 ));
892 }
893 let mut down_out = Vec::with_capacity(nl);
894 let mut gate_up = Vec::with_capacity(nl);
895 let mut kept_per_layer = Vec::with_capacity(nl);
896 for li in 0..nl {
897 let alive = &keep[li];
898 kept_per_layer.push(alive.iter().filter(|&&a| a).count());
899 let l = &fm.layers[li];
900 let mut down = match &ffn[li] {
901 Some((_, _, d)) => d.clone(),
902 None => l.down.clone(),
903 };
904 let hsz = fm.hidden;
905 for r in 0..hsz {
906 for (c, &a) in alive.iter().enumerate() {
907 if !a {
908 down[r * inter + c] = 0.0;
909 }
910 }
911 }
912 gate_up.push(ffn[li].as_ref().map(|(g, u, _)| (g.clone(), u.clone())));
913 down_out.push(down);
914 }
915 let total: usize = keep_visits
916 .iter()
917 .map(|a| a.iter().filter(|&&x| x).count())
918 .sum();
919 let report = BakeReport {
920 backbone,
921 masked,
922 overlaid,
923 pruned_ratio: 1.0 - total as f64 / (vn * inter) as f64,
924 kept_per_layer,
925 sec: t0.elapsed().as_secs_f64(),
926 };
927 let arts = BakeArtifacts {
928 keep,
929 keep_visits,
930 down: down_out,
931 gate_up,
932 fcd_layers: fcd,
933 };
934 Ok((report, arts))
935}
936
937fn keep_masks(logits: &[Vec<f32>], tau: f32, align: usize, uniform: bool) -> Vec<Vec<bool>> {
945 let inter = logits[0].len();
946 let round = |n: usize| -> usize {
947 let n = n.max(1);
948 if align <= 1 {
949 n.min(inter)
950 } else {
951 (n.div_ceil(align) * align).min(inter)
952 }
953 };
954 let mut want: Vec<usize> = logits
955 .iter()
956 .map(|l| round(l.iter().filter(|&&x| sigmoid(x) > tau).count()))
957 .collect();
958 if uniform {
959 let k = want.iter().copied().max().unwrap_or(inter);
960 want = vec![k; logits.len()];
961 }
962 logits
963 .iter()
964 .zip(&want)
965 .map(|(l, &k)| {
966 let mut idx: Vec<usize> = (0..inter).collect();
967 idx.sort_unstable_by(|&a, &b| l[b].total_cmp(&l[a]));
968 let mut alive = vec![false; inter];
969 for &i in idx.iter().take(k) {
970 alive[i] = true;
971 }
972 alive
973 })
974 .collect()
975}
976
977#[cfg(test)]
978mod tests {
979 use super::*;
980
981 fn kept(masks: &[Vec<bool>]) -> Vec<usize> {
982 masks
983 .iter()
984 .map(|m| m.iter().filter(|&&a| a).count())
985 .collect()
986 }
987
988 #[test]
991 fn keep_masks_aligns_up_and_preserves_alive() {
992 let inter = 96;
993 let l0: Vec<f32> = (0..inter)
996 .map(|i| if i < 40 { 1.0 } else { -1.0 - i as f32 * 0.01 })
997 .collect();
998 let l1: Vec<f32> = (0..inter)
1000 .map(|i| if i < 64 { 2.0 } else { -3.0 })
1001 .collect();
1002 let masks = keep_masks(&[l0.clone(), l1], 0.5, 32, false);
1003 assert_eq!(kept(&masks), vec![64, 64]);
1004 for i in 0..64 {
1007 assert!(masks[0][i], "neuron {i} should be kept");
1008 }
1009 for i in 64..inter {
1010 assert!(!masks[0][i], "neuron {i} should stay pruned");
1011 }
1012 }
1013
1014 #[test]
1016 fn keep_masks_uniform_takes_max() {
1017 let inter = 96;
1018 let l0: Vec<f32> = (0..inter)
1019 .map(|i| if i < 10 { 1.0 } else { -2.0 })
1020 .collect();
1021 let l1: Vec<f32> = (0..inter)
1022 .map(|i| if i < 70 { 1.0 } else { -2.0 })
1023 .collect();
1024 let masks = keep_masks(&[l0, l1], 0.5, 32, true);
1025 assert_eq!(kept(&masks), vec![96, 96]);
1026 }
1027
1028 #[test]
1031 fn keep_masks_edges() {
1032 let inter = 48;
1033 let l: Vec<f32> = (0..inter)
1034 .map(|i| if i < 47 { 1.0 } else { -2.0 })
1035 .collect();
1036 let masks = keep_masks(&[l.clone()], 0.5, 32, false);
1037 assert_eq!(kept(&masks), vec![48]); let masks = keep_masks(&[l], 0.5, 1, false);
1039 assert_eq!(kept(&masks), vec![47]);
1040 let dead: Vec<f32> = vec![-5.0; inter];
1041 let masks = keep_masks(&[dead], 0.5, 32, false);
1042 assert_eq!(kept(&masks), vec![32]); }
1044}