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, mats: &'b crate::fcd::LayerMats) -> 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: &[],
234 up: &[],
235 down: &mats.down,
236 gu: Some(&mats.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 mats_hold = fm.mats(li).expect("layer mats");
309 let wts = self.wts(li, &mats_hold);
310 let want = grad.is_some();
311 let (h2, a) = fm.layer_forward_scaled(li, &h, b, t, &wts, false, want, Some(&g));
312 h_ins.push(if want { h } else { Vec::new() });
313 acts.push(a);
314 masks.push(g);
315 h = h2;
316 if fm.loop_norm && li + 1 == nl && vl + 1 < vn {
319 let mut hn = vec![0f32; n * hsz];
320 let mut inv = vec![0f32; n];
321 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
322 if want {
323 lnorms[vl] = Some((h, inv));
324 }
325 h = hn;
326 }
327 }
328 let mut hn = vec![0f32; n * hsz];
330 let mut inv = vec![0f32; n];
331 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
332 let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
333 let vocab = lm.len() / hsz;
334 let pool = fm.pool.as_deref();
335 let mut nll = 0f64;
336 let mut dh_n = vec![0f32; n * hsz]; const POS_CHUNK: usize = 64;
341 let scored = b * (t - 1);
342 for bi in 0..b {
343 let base = bi * t;
344 let mut p0 = 0usize;
345 while p0 < t - 1 {
346 let pc = POS_CHUNK.min(t - 1 - p0);
347 let mut logits = vec![0f32; pc * vocab];
348 ops::gemm_nt(
349 &hn[(base + p0) * hsz..(base + p0 + pc) * hsz],
350 lm,
351 &mut logits,
352 pc,
353 hsz,
354 vocab,
355 pool,
356 );
357 for r in 0..pc {
358 let target = ids[base + p0 + r + 1] as usize;
359 let row = &mut logits[r * vocab..(r + 1) * vocab];
360 let mx = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
361 let mut sum = 0f64;
362 for v in row.iter() {
363 sum += ((*v as f64) - mx).exp();
364 }
365 nll += mx + sum.ln() - row[target] as f64;
366 if grad.is_some() {
367 let inv_n = 1.0 / scored as f64;
369 for v in row.iter_mut() {
370 *v = ((((*v as f64) - mx).exp() / sum) * inv_n) as f32;
371 }
372 row[target] -= inv_n as f32;
373 }
374 }
375 if grad.is_some() {
376 ops::gemm_dx(
377 &logits,
378 lm,
379 &mut dh_n[(base + p0) * hsz..(base + p0 + pc) * hsz],
380 pc,
381 hsz,
382 vocab,
383 pool,
384 );
385 }
386 p0 += pc;
387 }
388 }
389 let Some((dmask, dffn)) = grad else {
390 return (nll, scored);
391 };
392 let t_bwd = std::time::Instant::now();
394 let mut dh = vec![0f32; n * hsz];
395 ops::rmsnorm_bwd(&h, &fm.final_norm, &inv, &dh_n, fm.gemma, &mut dh, None);
396 for vl in (0..vn).rev() {
397 let li = vl % nl;
398 if let Some((hb, inv)) = lnorms[vl].as_ref() {
400 let mut dprev = vec![0f32; n * hsz];
401 ops::rmsnorm_bwd(hb, &fm.final_norm, inv, &dh, fm.gemma, &mut dprev, None);
402 dh = dprev;
403 }
404 let a = acts[vl].as_ref().expect("acts saved in grad mode");
405 let g = &masks[vl];
406 let inter = fm.layers[li].inter;
407 let mats_hold = fm.mats(li).expect("layer mats");
408 let wts = self.wts(li, &mats_hold);
409 let mut dact2 = vec![0f32; t * inter];
411 ops::gemm_dx(&dh, wts.down, &mut dact2, t, inter, hsz, fm.pool.as_deref());
412 if let Some((_, _, dd)) = dffn[li].as_mut() {
413 let mut act2 = a.act.clone();
415 for r in 0..t {
416 for (x, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
417 *x *= gv;
418 }
419 }
420 let mut dw = vec![0f32; hsz * inter];
421 ops::gemm_dw(&dh, &act2, &mut dw, t, inter, hsz, fm.pool.as_deref());
422 for (o, &x) in dd.iter_mut().zip(&dw) {
423 *o += x as f64;
424 }
425 }
426 {
430 let dm = &mut dmask[vl];
431 for r in 0..t {
432 let da = &dact2[r * inter..(r + 1) * inter];
433 let aa = &a.act[r * inter..(r + 1) * inter];
434 for j in 0..inter {
435 dm[j] += da[j] as f64 * aa[j] as f64;
436 }
437 }
438 for (j, d) in dm.iter_mut().enumerate() {
440 let _ = j;
441 let _ = d;
442 }
443 }
444 let mut dg_pre = vec![0f32; t * inter];
446 let mut du_pre = vec![0f32; t * inter];
447 for r in 0..t {
448 for j in 0..inter {
449 let i = r * inter + j;
450 let da = dact2[i] * g[j];
451 let sg = ops::silu(a.gpre[i]);
452 dg_pre[i] = da * a.upre[i] * ops::silu_bwd(a.gpre[i]);
453 du_pre[i] = da * sg;
454 }
455 }
456 let mut dn2 = vec![0f32; t * hsz];
459 if let Some(gu) = wts.gu {
460 let mut dgu = vec![0f32; t * 2 * inter];
461 for r in 0..t {
462 let row = &mut dgu[r * 2 * inter..(r + 1) * 2 * inter];
463 row[..inter].copy_from_slice(&dg_pre[r * inter..(r + 1) * inter]);
464 row[inter..].copy_from_slice(&du_pre[r * inter..(r + 1) * inter]);
465 }
466 ops::gemm_dx(&dgu, gu, &mut dn2, t, hsz, 2 * inter, fm.pool.as_deref());
467 } else {
468 ops::gemm_dx(
469 &dg_pre,
470 wts.gate,
471 &mut dn2,
472 t,
473 hsz,
474 inter,
475 fm.pool.as_deref(),
476 );
477 let mut dn2b = vec![0f32; t * hsz];
478 ops::gemm_dx(
479 &du_pre,
480 wts.up,
481 &mut dn2b,
482 t,
483 hsz,
484 inter,
485 fm.pool.as_deref(),
486 );
487 for (x, &y) in dn2.iter_mut().zip(&dn2b) {
488 *x += y;
489 }
490 }
491 if let Some((dgw, duw, _)) = dffn[li].as_mut() {
492 let mut dw = vec![0f32; inter * hsz];
493 ops::gemm_dw(&dg_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
494 for (o, &x) in dgw.iter_mut().zip(&dw) {
495 *o += x as f64;
496 }
497 dw.fill(0.0);
498 ops::gemm_dw(&du_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
499 for (o, &x) in duw.iter_mut().zip(&dw) {
500 *o += x as f64;
501 }
502 }
503 let mut dh1 = dh.clone(); ops::rmsnorm_bwd(&a.h1, wts.pln, &a.inv2, &dn2, fm.gemma, &mut dh1, None);
507 dh = dh1;
508 let _ = &h_ins[vl];
509 }
510 crate::fcd::prof::add(&crate::fcd::prof::BWD, t_bwd);
511 (nll, scored)
512 }
513}
514
515fn held_ppl(pass: &Pass, held: &[Vec<u32>]) -> f64 {
517 if held.is_empty() {
521 return f64::NAN;
522 }
523 let t = held[0].len();
524 if held.iter().all(|c| c.len() == t) {
525 let flat: Vec<u32> = held.iter().flatten().copied().collect();
526 let (l, k) = pass.chunk_batch(&flat, held.len(), None);
527 return (l / k.max(1) as f64).exp();
528 }
529 let mut nll = 0f64;
530 let mut n = 0usize;
531 for c in held {
532 let (l, k) = pass.chunk(c, None);
533 nll += l;
534 n += k;
535 }
536 (nll / n.max(1) as f64).exp()
537}
538
539pub fn replica_score_file_mask(
545 model: &Arc<CmfModel>,
546 chunks: &[Vec<u32>],
547) -> Result<(f64, f64), String> {
548 let o1_off = crate::nystrom::O1Cfg {
549 layers: crate::nystrom::O1Layers::List(Vec::new()),
550 m: 4,
551 w: 8,
552 sink: 1,
553 rect: crate::nystrom::O1_DEFAULT_RECT,
554 };
555 let fm = FcdModel::from_cmf(model, &o1_off)?;
556 let nl = fm.layers.len();
557 let loops = fm.loops.max(1);
558 let vn = nl * loops;
559 let inter = fm.layers[0].inter;
560 let ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
561 let task = &model.masks.default_task;
563 let mask = model
564 .masks
565 .masks
566 .iter()
567 .find(|m| &m.name == task)
568 .or_else(|| model.masks.masks.first());
569 let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
570 let masked_logits: Vec<Vec<f32>> = match mask {
571 Some(m) => (0..vn)
572 .map(|vl| {
573 let row = m.ffn_masks.get(vl).map(|v| v.as_slice()).unwrap_or(&[]);
574 (0..inter)
575 .map(|j| {
576 if (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0 {
577 50.0
578 } else {
579 -50.0
580 }
581 })
582 .collect()
583 })
584 .collect(),
585 None => open.clone(),
586 };
587 let score = |logits: &[Vec<f32>]| -> f64 {
588 let pass = Pass {
589 fm: &fm,
590 tau: 0.5,
591 logits,
592 hard: true,
593 ffn: &ffn,
594 };
595 held_ppl(&pass, chunks)
596 };
597 Ok((score(&open), score(&masked_logits)))
598}
599
600pub fn skill_bake(
602 model: &Arc<CmfModel>,
603 chunks: &[Vec<u32>],
604 held_n: usize,
605 hy: &BakeHyper,
606 mut log: impl FnMut(&str),
607) -> Result<(BakeReport, BakeArtifacts), String> {
608 let t0 = std::time::Instant::now();
609 let o1_off = crate::nystrom::O1Cfg {
610 layers: crate::nystrom::O1Layers::List(Vec::new()),
611 m: 4,
612 w: 8,
613 sink: 1,
614 rect: crate::nystrom::O1_DEFAULT_RECT,
615 };
616 let fm = FcdModel::from_cmf(model, &o1_off)?;
617 let nl = fm.layers.len();
618 let inter = fm.layers.iter().map(|l| l.inter).max().unwrap_or(0);
619 if fm.layers.iter().any(|l| l.inter != inter) {
620 return Err("skill bake: non-uniform FFN widths".into());
621 }
622 let held: Vec<Vec<u32>> = chunks[..held_n.min(chunks.len())].to_vec();
623 let calib: Vec<Vec<u32>> = chunks[held_n.min(chunks.len())..].to_vec();
624 if calib.len() < 12 {
625 return Err(format!(
626 "skill bake: corpus too small ({} calib chunks)",
627 calib.len()
628 ));
629 }
630 let fcd: Vec<usize> = (nl.saturating_sub(hy.fcd_layers)..nl).collect();
631 let _rng = SplitMix64::new(hy.seed);
632
633 let loops = fm.loops.max(1);
661 let m0 = mask_init_logit(loops);
662 let vn = nl * loops;
664 let mut logits: Vec<Vec<f32>> = vec![vec![m0; inter]; vn];
665 let mut ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
666
667 let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
670 let base_pass = Pass {
671 fm: &fm,
672 tau: hy.tau,
673 logits: &open,
674 hard: true,
675 ffn: &ffn,
676 };
677 let backbone = held_ppl(&base_pass, &held);
678 log(&format!("baseline (full): {backbone:.3}"));
679
680 let mut adam_a = Adam::new(&vec![inter; vn], hy.lr_a);
682 let mut l1 = hy.l1_init * hy.l1_mult;
683 let l1_step_eff = hy.l1_step * hy.l1_mult;
684 let mut best: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
686 let mut max_sp: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
688 let mut prev_alive: Option<Vec<Vec<bool>>> = None;
689 let mut acc_chunk = 0f64;
693 let mut acc_adam = 0f64;
694 crate::gpu::bake_precision_strict(true);
702 for step in 0..hy.steps_a {
703 let t_step = std::time::Instant::now();
704 let chunk = &calib[step % calib.len()];
705 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
706 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = vec![None; nl];
707 let pass = Pass {
708 fm: &fm,
709 tau: hy.tau,
710 logits: &logits,
711 hard: false,
712 ffn: &ffn,
713 };
714 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
715 let l1_per = l1 / (inter as f64 * nl as f64);
717 for li in 0..vn {
718 for j in 0..inter {
719 let s = sigmoid(logits[li][j]) as f64;
720 dmask[li][j] = dmask[li][j] * s * (1.0 - s) + l1_per * s * (1.0 - s);
721 }
722 }
723 let t_chunk = t_step.elapsed().as_secs_f64();
736 let mut params: Vec<&mut [f32]> = logits.iter_mut().map(|v| v.as_mut_slice()).collect();
737 adam_a.step(&mut params, &dmask, 1.0);
738 acc_chunk += t_chunk;
739 acc_adam += t_step.elapsed().as_secs_f64() - t_chunk;
740 if (step + 1) % hy.eval_every == 0 {
741 l1 += l1_step_eff;
742 let pass = Pass {
743 fm: &fm,
744 tau: hy.tau,
745 logits: &logits,
746 hard: true,
747 ffn: &ffn,
748 };
749 crate::gpu::bake_precision_strict(false);
750 let hp = held_ppl(&pass, &held);
751 crate::gpu::bake_precision_strict(true);
752 let cur: Vec<Vec<bool>> = logits
758 .iter()
759 .map(|l| l.iter().map(|&x| sigmoid(x) > hy.tau).collect())
760 .collect();
761 if let Some(prev) = &prev_alive {
762 let died: Vec<String> = cur
763 .iter()
764 .zip(prev)
765 .enumerate()
766 .flat_map(|(li, (c, p))| {
767 c.iter()
768 .zip(p.iter())
769 .enumerate()
770 .filter(|&(_, (&cj, &pj))| pj && !cj)
771 .map(move |(j, _)| format!("L{li}:{j}"))
772 })
773 .collect();
774 if !died.is_empty() {
775 log(&format!(
776 " closed since last eval: {}: {}{}",
777 died.len(),
778 died.iter().take(32).cloned().collect::<Vec<_>>().join(" "),
779 if died.len() > 32 { " …" } else { "" }
780 ));
781 }
782 }
783 let alive: usize = cur.iter().map(|l| l.iter().filter(|&&b| b).count()).sum();
784 prev_alive = Some(cur);
785 let sp = 1.0 - alive as f64 / (vn * inter) as f64;
786 if sp > max_sp.2 {
788 max_sp = (hp, Some(logits.clone()), sp);
789 }
790 if hy.target_sparsity > 0.0 {
792 if sp >= hy.target_sparsity && hp < best.0 {
793 best = (hp, Some(logits.clone()), sp);
794 }
795 } else if hp < best.0 {
796 best = (hp, Some(logits.clone()), sp);
797 }
798 log(&format!(
799 " [A] step {}: L1={l1:.3} pruned={:.2}% hard-PPL={hp:.3} (bottom {}@{:.2}%) [fwd+bwd {:.1}s, adam {:.2}s per step]",
800 step + 1,
801 sp * 100.0,
802 if best.0 == f64::MAX {
803 "—".to_string()
804 } else {
805 format!("{:.3}", best.0)
806 },
807 best.2 * 100.0,
808 acc_chunk / (step + 1) as f64,
809 acc_adam / (step + 1) as f64
810 ));
811 }
812 }
813 crate::gpu::bake_precision_strict(false);
817 if hy.target_sparsity > 0.0 && best.1.is_none() {
818 log(&format!(
819 "[A] target sparsity {:.0}% not reached; using max-sparsity checkpoint ({:.0}%)",
820 hy.target_sparsity * 100.0,
821 max_sp.2 * 100.0
822 ));
823 best = max_sp;
824 }
825 {
828 use crate::fcd::prof;
829 let (a, f, bw, g, gc) = (
830 prof::take(&prof::ATTN_FWD),
831 prof::take(&prof::FFN_FWD),
832 prof::take(&prof::BWD),
833 prof::take(&prof::GEMM),
834 prof::GEMM_CALLS.swap(0, std::sync::atomic::Ordering::Relaxed),
835 );
836 log(&format!(
837 "[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)",
838 hy.steps_a,
839 if gc > 0 { g * 1000.0 / gc as f64 } else { 0.0 }
840 ));
841 log(&format!("[prof] gemm shapes:\n{}", prof::shape_report(6)));
842 }
843 if let Some(b) = best.1.take() {
844 logits = b;
845 }
846 let pass = Pass {
847 fm: &fm,
848 tau: hy.tau,
849 logits: &logits,
850 hard: true,
851 ffn: &ffn,
852 };
853 let masked = held_ppl(&pass, &held);
854 log(&format!(
855 "[A] {:.0}s: masked-PPL {masked:.3}",
856 t0.elapsed().as_secs_f64()
857 ));
858
859 for &li in &fcd {
861 let p = format!("model.layers.{li}.");
862 ffn[li] = Some((
863 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.gate_proj.weight"))
864 .map_err(|e| format!("phase-B gate: {e}"))?,
865 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.up_proj.weight"))
866 .map_err(|e| format!("phase-B up: {e}"))?,
867 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.down_proj.weight"))
868 .map_err(|e| format!("phase-B down: {e}"))?,
869 ));
870 }
871 let sizes: Vec<usize> = fcd
872 .iter()
873 .flat_map(|&li| {
874 let (g, u, d) = ffn[li].as_ref().expect("phase-B masters");
875 [g.len(), u.len(), d.len()]
876 })
877 .collect();
878 let mut adam_b = Adam::new(&sizes, hy.lr_b);
879 let mut best_b: (f64, Option<Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>>>) = (masked, None);
880 for step in 0..hy.steps_b {
881 let chunk = &calib[step % calib.len()];
882 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
883 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = (0..nl)
884 .map(|li| {
885 ffn[li]
886 .as_ref()
887 .map(|(g, u, d)| (vec![0.0; g.len()], vec![0.0; u.len()], vec![0.0; d.len()]))
888 })
889 .collect();
890 let pass = Pass {
891 fm: &fm,
892 tau: hy.tau,
893 logits: &logits,
894 hard: true,
895 ffn: &ffn,
896 };
897 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
898 let lr_scale = 0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
900 let first_fcd = fcd[0];
901 let mut params: Vec<&mut [f32]> = Vec::new();
902 let mut grads: Vec<Vec<f64>> = Vec::new();
903 for (off, slot) in ffn[first_fcd..].iter_mut().enumerate() {
904 let li = first_fcd + off;
905 let Some((g, u, d)) = slot.as_mut() else {
906 continue;
907 };
908 let (dg, du, dd) = dffn[li].take().unwrap();
909 params.push(g.as_mut_slice());
910 grads.push(dg);
911 params.push(u.as_mut_slice());
912 grads.push(du);
913 params.push(d.as_mut_slice());
914 grads.push(dd);
915 }
916 adam_b.step(&mut params, &grads, lr_scale * mask_step_scale(loops));
921 if (step + 1) % hy.eval_every == 0 {
922 let pass = Pass {
923 fm: &fm,
924 tau: hy.tau,
925 logits: &logits,
926 hard: true,
927 ffn: &ffn,
928 };
929 let cur = held_ppl(&pass, &held);
930 if cur < best_b.0 {
931 best_b = (cur, Some(ffn.clone()));
932 }
933 log(&format!(
934 " [B] step {}: held-PPL {cur:.3} (best {:.3})",
935 step + 1,
936 best_b.0
937 ));
938 }
939 }
940 if let Some(b) = best_b.1.take() {
941 ffn = b;
942 }
943 let overlaid = best_b.0;
944
945 let keep_visits = keep_masks(&logits, hy.tau, hy.align, hy.uniform_inter);
950 let keep: Vec<Vec<bool>> = (0..nl)
951 .map(|li| {
952 (0..inter)
953 .map(|j| (0..loops).any(|v| keep_visits[v * nl + li][j]))
954 .collect()
955 })
956 .collect();
957 if hy.align > 1 || hy.uniform_inter {
958 let raw: usize = logits
959 .iter()
960 .map(|l| l.iter().filter(|&&x| sigmoid(x) > hy.tau).count())
961 .sum();
962 let padded: usize = keep_visits
966 .iter()
967 .map(|a| a.iter().filter(|&&x| x).count())
968 .sum::<usize>()
969 .saturating_sub(raw);
970 log(&format!(
971 "align: +{padded} neurons resurrected (align {}, uniform {})",
972 hy.align, hy.uniform_inter
973 ));
974 }
975 let mut down_out = Vec::with_capacity(nl);
976 let mut gate_up = Vec::with_capacity(nl);
977 let mut kept_per_layer = Vec::with_capacity(nl);
978 for li in 0..nl {
979 let alive = &keep[li];
980 kept_per_layer.push(alive.iter().filter(|&&a| a).count());
981 let mut down = match &ffn[li] {
982 Some((_, _, d)) => d.clone(),
983 None => fm.mats(li).expect("layer mats").down.clone(),
984 };
985 let hsz = fm.hidden;
986 for r in 0..hsz {
987 for (c, &a) in alive.iter().enumerate() {
988 if !a {
989 down[r * inter + c] = 0.0;
990 }
991 }
992 }
993 gate_up.push(ffn[li].as_ref().map(|(g, u, _)| (g.clone(), u.clone())));
994 down_out.push(down);
995 }
996 let total: usize = keep_visits
997 .iter()
998 .map(|a| a.iter().filter(|&&x| x).count())
999 .sum();
1000 let report = BakeReport {
1001 backbone,
1002 masked,
1003 overlaid,
1004 pruned_ratio: 1.0 - total as f64 / (vn * inter) as f64,
1005 kept_per_layer,
1006 sec: t0.elapsed().as_secs_f64(),
1007 };
1008 let arts = BakeArtifacts {
1009 keep,
1010 keep_visits,
1011 down: down_out,
1012 gate_up,
1013 fcd_layers: fcd,
1014 };
1015 Ok((report, arts))
1016}
1017
1018fn keep_masks(logits: &[Vec<f32>], tau: f32, align: usize, uniform: bool) -> Vec<Vec<bool>> {
1026 let inter = logits[0].len();
1027 let round = |n: usize| -> usize {
1028 let n = n.max(1);
1029 if align <= 1 {
1030 n.min(inter)
1031 } else {
1032 (n.div_ceil(align) * align).min(inter)
1033 }
1034 };
1035 let mut want: Vec<usize> = logits
1036 .iter()
1037 .map(|l| round(l.iter().filter(|&&x| sigmoid(x) > tau).count()))
1038 .collect();
1039 if uniform {
1040 let k = want.iter().copied().max().unwrap_or(inter);
1041 want = vec![k; logits.len()];
1042 }
1043 logits
1044 .iter()
1045 .zip(&want)
1046 .map(|(l, &k)| {
1047 let mut idx: Vec<usize> = (0..inter).collect();
1048 idx.sort_unstable_by(|&a, &b| l[b].total_cmp(&l[a]));
1049 let mut alive = vec![false; inter];
1050 for &i in idx.iter().take(k) {
1051 alive[i] = true;
1052 }
1053 alive
1054 })
1055 .collect()
1056}
1057
1058#[cfg(test)]
1059mod tests {
1060 use super::*;
1061
1062 fn kept(masks: &[Vec<bool>]) -> Vec<usize> {
1063 masks
1064 .iter()
1065 .map(|m| m.iter().filter(|&&a| a).count())
1066 .collect()
1067 }
1068
1069 #[test]
1072 fn keep_masks_aligns_up_and_preserves_alive() {
1073 let inter = 96;
1074 let l0: Vec<f32> = (0..inter)
1077 .map(|i| if i < 40 { 1.0 } else { -1.0 - i as f32 * 0.01 })
1078 .collect();
1079 let l1: Vec<f32> = (0..inter)
1081 .map(|i| if i < 64 { 2.0 } else { -3.0 })
1082 .collect();
1083 let masks = keep_masks(&[l0.clone(), l1], 0.5, 32, false);
1084 assert_eq!(kept(&masks), vec![64, 64]);
1085 for i in 0..64 {
1088 assert!(masks[0][i], "neuron {i} should be kept");
1089 }
1090 for i in 64..inter {
1091 assert!(!masks[0][i], "neuron {i} should stay pruned");
1092 }
1093 }
1094
1095 #[test]
1097 fn keep_masks_uniform_takes_max() {
1098 let inter = 96;
1099 let l0: Vec<f32> = (0..inter)
1100 .map(|i| if i < 10 { 1.0 } else { -2.0 })
1101 .collect();
1102 let l1: Vec<f32> = (0..inter)
1103 .map(|i| if i < 70 { 1.0 } else { -2.0 })
1104 .collect();
1105 let masks = keep_masks(&[l0, l1], 0.5, 32, true);
1106 assert_eq!(kept(&masks), vec![96, 96]);
1107 }
1108
1109 #[test]
1112 fn keep_masks_edges() {
1113 let inter = 48;
1114 let l: Vec<f32> = (0..inter)
1115 .map(|i| if i < 47 { 1.0 } else { -2.0 })
1116 .collect();
1117 let masks = keep_masks(&[l.clone()], 0.5, 32, false);
1118 assert_eq!(kept(&masks), vec![48]); let masks = keep_masks(&[l], 0.5, 1, false);
1120 assert_eq!(kept(&masks), vec![47]);
1121 let dead: Vec<f32> = vec![-5.0; inter];
1122 let masks = keep_masks(&[dead], 0.5, 32, false);
1123 assert_eq!(kept(&masks), vec![32]); }
1125}