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 pub logits: Vec<Vec<f32>>,
110}
111
112const CLIP: f64 = 1.0;
113const B1: f64 = 0.9;
114const B2: f64 = 0.999;
115const EPS: f64 = 1e-8;
116
117struct Adam {
119 m: Vec<Vec<f64>>,
120 v: Vec<Vec<f64>>,
121 t: i32,
122 lr: f64,
123}
124
125impl Adam {
126 fn new(sizes: &[usize], lr: f64) -> Self {
127 Self {
128 m: sizes.iter().map(|&n| vec![0.0; n]).collect(),
129 v: sizes.iter().map(|&n| vec![0.0; n]).collect(),
130 t: 0,
131 lr,
132 }
133 }
134
135 fn step(&mut self, params: &mut [&mut [f32]], grads: &[Vec<f64>], lr_scale: f64) {
137 let gn: f64 = grads
138 .iter()
139 .flat_map(|g| g.iter().map(|x| x * x))
140 .sum::<f64>()
141 .sqrt();
142 let clip = if gn > CLIP { CLIP / gn } else { 1.0 };
143 self.t += 1;
144 let (bc1, bc2) = (1.0 - B1.powi(self.t), 1.0 - B2.powi(self.t));
145 for (pi, p) in params.iter_mut().enumerate() {
146 for j in 0..p.len() {
147 let g = grads[pi][j] * clip;
148 let m = &mut self.m[pi][j];
149 let v = &mut self.v[pi][j];
150 *m = B1 * *m + (1.0 - B1) * g;
151 *v = B2 * *v + (1.0 - B2) * g * g;
152 let upd = (*m / bc1) / ((*v / bc2).sqrt() + EPS);
153 p[j] -= (self.lr * lr_scale * upd) as f32;
154 }
155 }
156 }
157}
158
159pub fn mask_init_logit(loops: usize) -> f32 {
178 let base = 1.0f32 / (1.0 + (-2.0f32).exp());
179 let per_visit = base.powf(1.0 / loops.max(1) as f32);
180 (per_visit / (1.0 - per_visit)).ln()
181}
182
183pub fn mask_step_scale(loops: usize) -> f64 {
189 1.0 / loops.max(1) as f64
190}
191
192fn sigmoid(x: f32) -> f32 {
193 1.0 / (1.0 + (-x).exp())
194}
195
196struct Pass<'a> {
199 fm: &'a FcdModel,
200 tau: f32,
201 logits: &'a [Vec<f32>],
203 hard: bool,
204 ffn: &'a [Option<(Vec<f32>, Vec<f32>, Vec<f32>)>],
206}
207
208impl Pass<'_> {
209 fn gates(&self, li: usize) -> Vec<f32> {
210 self.logits[li]
211 .iter()
212 .map(|&l| {
213 let s = sigmoid(l);
214 if self.hard {
215 if s > self.tau { 1.0 } else { 0.0 }
216 } else {
217 s
218 }
219 })
220 .collect()
221 }
222
223 fn wts<'b>(&'b self, li: usize, mats: &'b crate::fcd::LayerMats) -> LnFfn<'b> {
224 let l = &self.fm.layers[li];
225 match &self.ffn[li] {
226 Some((g, u, d)) => LnFfn {
227 iln: &l.iln,
228 pln: &l.pln,
229 gate: g,
230 up: u,
231 down: d,
232 gu: None,
234 },
235 None => LnFfn {
236 iln: &l.iln,
237 pln: &l.pln,
238 gate: &[],
239 up: &[],
240 down: &mats.down,
241 gu: Some(&mats.gu),
242 },
243 }
244 }
245
246 #[allow(clippy::too_many_arguments)]
250 fn chunk(
251 &self,
252 ids: &[u32],
253 grad: Option<(
254 &mut [Vec<f64>],
255 &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
256 )>,
257 ) -> (f64, usize) {
258 self.chunk_batch(ids, 1, grad)
259 }
260
261 fn chunk_batch(
269 &self,
270 ids: &[u32],
271 b: usize,
272 grad: Option<(
273 &mut [Vec<f64>],
274 &mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
275 )>,
276 ) -> (f64, usize) {
277 let fm = self.fm;
278 let hsz = fm.hidden;
279 debug_assert!(ids.len() % b.max(1) == 0, "ragged batch");
280 debug_assert!(grad.is_none() || b == 1, "grads are per-chunk");
283 let t = ids.len() / b.max(1);
284 let n = b * t;
285 let nl = fm.layers.len();
286 let mut h = vec![0f32; n * hsz];
288 for (r, &id) in ids.iter().enumerate() {
289 h[r * hsz..(r + 1) * hsz]
290 .copy_from_slice(&fm.embed[id as usize * hsz..(id as usize + 1) * hsz]);
291 }
292 let loops = fm.loops.max(1);
300 let vn = nl * loops;
301 let mut h_ins = Vec::with_capacity(vn);
302 let mut acts = Vec::with_capacity(vn);
303 let mut masks = Vec::with_capacity(vn);
304 let mut lnorms: Vec<Option<(Vec<f32>, Vec<f32>)>> = vec![None; vn];
306 for vl in 0..vn {
307 let li = vl % nl;
308 let g = self.gates(vl);
313 let mats_hold = fm.mats(li).expect("layer mats");
314 let wts = self.wts(li, &mats_hold);
315 let want = grad.is_some();
316 let (h2, a) = fm.layer_forward_scaled(li, &h, b, t, &wts, false, want, Some(&g));
317 h_ins.push(if want { h } else { Vec::new() });
318 acts.push(a);
319 masks.push(g);
320 h = h2;
321 if fm.loop_norm && li + 1 == nl && vl + 1 < vn {
324 let mut hn = vec![0f32; n * hsz];
325 let mut inv = vec![0f32; n];
326 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
327 if want {
328 lnorms[vl] = Some((h, inv));
329 }
330 h = hn;
331 }
332 }
333 let mut hn = vec![0f32; n * hsz];
335 let mut inv = vec![0f32; n];
336 ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
337 let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
338 let vocab = lm.len() / hsz;
339 let pool = fm.pool.as_deref();
340 let mut nll = 0f64;
341 let mut dh_n = vec![0f32; n * hsz]; const POS_CHUNK: usize = 64;
346 let scored = b * (t - 1);
347 for bi in 0..b {
348 let base = bi * t;
349 let mut p0 = 0usize;
350 while p0 < t - 1 {
351 let pc = POS_CHUNK.min(t - 1 - p0);
352 let mut logits = vec![0f32; pc * vocab];
353 ops::gemm_nt(
354 &hn[(base + p0) * hsz..(base + p0 + pc) * hsz],
355 lm,
356 &mut logits,
357 pc,
358 hsz,
359 vocab,
360 pool,
361 );
362 for r in 0..pc {
363 let target = ids[base + p0 + r + 1] as usize;
364 let row = &mut logits[r * vocab..(r + 1) * vocab];
365 let mx = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
366 let mut sum = 0f64;
367 for v in row.iter() {
368 sum += ((*v as f64) - mx).exp();
369 }
370 nll += mx + sum.ln() - row[target] as f64;
371 if grad.is_some() {
372 let inv_n = 1.0 / scored as f64;
374 for v in row.iter_mut() {
375 *v = ((((*v as f64) - mx).exp() / sum) * inv_n) as f32;
376 }
377 row[target] -= inv_n as f32;
378 }
379 }
380 if grad.is_some() {
381 ops::gemm_dx(
382 &logits,
383 lm,
384 &mut dh_n[(base + p0) * hsz..(base + p0 + pc) * hsz],
385 pc,
386 hsz,
387 vocab,
388 pool,
389 );
390 }
391 p0 += pc;
392 }
393 }
394 let Some((dmask, dffn)) = grad else {
395 return (nll, scored);
396 };
397 let t_bwd = std::time::Instant::now();
399 let mut dh = vec![0f32; n * hsz];
400 ops::rmsnorm_bwd(&h, &fm.final_norm, &inv, &dh_n, fm.gemma, &mut dh, None);
401 for vl in (0..vn).rev() {
402 let li = vl % nl;
403 if let Some((hb, inv)) = lnorms[vl].as_ref() {
405 let mut dprev = vec![0f32; n * hsz];
406 ops::rmsnorm_bwd(hb, &fm.final_norm, inv, &dh, fm.gemma, &mut dprev, None);
407 dh = dprev;
408 }
409 let a = acts[vl].as_ref().expect("acts saved in grad mode");
410 let g = &masks[vl];
411 let inter = fm.layers[li].inter;
412 let mats_hold = fm.mats(li).expect("layer mats");
413 let wts = self.wts(li, &mats_hold);
414 let mut dact2 = vec![0f32; t * inter];
416 ops::gemm_dx(&dh, wts.down, &mut dact2, t, inter, hsz, fm.pool.as_deref());
417 if let Some((_, _, dd)) = dffn[li].as_mut() {
418 let mut act2 = a.act.clone();
420 for r in 0..t {
421 for (x, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
422 *x *= gv;
423 }
424 }
425 let mut dw = vec![0f32; hsz * inter];
426 ops::gemm_dw(&dh, &act2, &mut dw, t, inter, hsz, fm.pool.as_deref());
427 for (o, &x) in dd.iter_mut().zip(&dw) {
428 *o += x as f64;
429 }
430 }
431 {
435 let dm = &mut dmask[vl];
436 for r in 0..t {
437 let da = &dact2[r * inter..(r + 1) * inter];
438 let aa = &a.act[r * inter..(r + 1) * inter];
439 for j in 0..inter {
440 dm[j] += da[j] as f64 * aa[j] as f64;
441 }
442 }
443 for (j, d) in dm.iter_mut().enumerate() {
445 let _ = j;
446 let _ = d;
447 }
448 }
449 let mut dg_pre = vec![0f32; t * inter];
451 let mut du_pre = vec![0f32; t * inter];
452 for r in 0..t {
453 for j in 0..inter {
454 let i = r * inter + j;
455 let da = dact2[i] * g[j];
456 let sg = ops::silu(a.gpre[i]);
457 dg_pre[i] = da * a.upre[i] * ops::silu_bwd(a.gpre[i]);
458 du_pre[i] = da * sg;
459 }
460 }
461 let mut dn2 = vec![0f32; t * hsz];
464 if let Some(gu) = wts.gu {
465 let mut dgu = vec![0f32; t * 2 * inter];
466 for r in 0..t {
467 let row = &mut dgu[r * 2 * inter..(r + 1) * 2 * inter];
468 row[..inter].copy_from_slice(&dg_pre[r * inter..(r + 1) * inter]);
469 row[inter..].copy_from_slice(&du_pre[r * inter..(r + 1) * inter]);
470 }
471 ops::gemm_dx(&dgu, gu, &mut dn2, t, hsz, 2 * inter, fm.pool.as_deref());
472 } else {
473 ops::gemm_dx(
474 &dg_pre,
475 wts.gate,
476 &mut dn2,
477 t,
478 hsz,
479 inter,
480 fm.pool.as_deref(),
481 );
482 let mut dn2b = vec![0f32; t * hsz];
483 ops::gemm_dx(
484 &du_pre,
485 wts.up,
486 &mut dn2b,
487 t,
488 hsz,
489 inter,
490 fm.pool.as_deref(),
491 );
492 for (x, &y) in dn2.iter_mut().zip(&dn2b) {
493 *x += y;
494 }
495 }
496 if let Some((dgw, duw, _)) = dffn[li].as_mut() {
497 let mut dw = vec![0f32; inter * hsz];
498 ops::gemm_dw(&dg_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
499 for (o, &x) in dgw.iter_mut().zip(&dw) {
500 *o += x as f64;
501 }
502 dw.fill(0.0);
503 ops::gemm_dw(&du_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
504 for (o, &x) in duw.iter_mut().zip(&dw) {
505 *o += x as f64;
506 }
507 }
508 let mut dh1 = dh.clone(); ops::rmsnorm_bwd(&a.h1, wts.pln, &a.inv2, &dn2, fm.gemma, &mut dh1, None);
512 dh = dh1;
513 let _ = &h_ins[vl];
514 }
515 crate::fcd::prof::add(&crate::fcd::prof::BWD, t_bwd);
516 (nll, scored)
517 }
518}
519
520fn held_ppl(pass: &Pass, held: &[Vec<u32>]) -> f64 {
522 if held.is_empty() {
526 return f64::NAN;
527 }
528 let t = held[0].len();
529 if held.iter().all(|c| c.len() == t) {
530 let flat: Vec<u32> = held.iter().flatten().copied().collect();
531 let (l, k) = pass.chunk_batch(&flat, held.len(), None);
532 return (l / k.max(1) as f64).exp();
533 }
534 let mut nll = 0f64;
535 let mut n = 0usize;
536 for c in held {
537 let (l, k) = pass.chunk(c, None);
538 nll += l;
539 n += k;
540 }
541 (nll / n.max(1) as f64).exp()
542}
543
544pub fn replica_score_file_mask(
550 model: &Arc<CmfModel>,
551 chunks: &[Vec<u32>],
552) -> Result<(f64, f64), String> {
553 let o1_off = crate::nystrom::O1Cfg {
554 layers: crate::nystrom::O1Layers::List(Vec::new()),
555 m: 4,
556 w: 8,
557 sink: 1,
558 rect: crate::nystrom::O1_DEFAULT_RECT,
559 };
560 let fm = FcdModel::from_cmf(model, &o1_off, false)?;
561 let nl = fm.layers.len();
562 let loops = fm.loops.max(1);
563 let vn = nl * loops;
564 let inter = fm.layers[0].inter;
565 let ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
566 let task = &model.masks.default_task;
568 let mask = model
569 .masks
570 .masks
571 .iter()
572 .find(|m| &m.name == task)
573 .or_else(|| model.masks.masks.first());
574 let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
575 let masked_logits: Vec<Vec<f32>> = match mask {
576 Some(m) => (0..vn)
577 .map(|vl| {
578 let row = m.ffn_masks.get(vl).map(|v| v.as_slice()).unwrap_or(&[]);
579 (0..inter)
580 .map(|j| {
581 if (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0 {
582 50.0
583 } else {
584 -50.0
585 }
586 })
587 .collect()
588 })
589 .collect(),
590 None => open.clone(),
591 };
592 let score = |logits: &[Vec<f32>]| -> f64 {
593 let pass = Pass {
594 fm: &fm,
595 tau: 0.5,
596 logits,
597 hard: true,
598 ffn: &ffn,
599 };
600 held_ppl(&pass, chunks)
601 };
602 Ok((score(&open), score(&masked_logits)))
603}
604
605pub fn skill_bake(
607 model: &Arc<CmfModel>,
608 chunks: &[Vec<u32>],
609 held_n: usize,
610 hy: &BakeHyper,
611 mut log: impl FnMut(&str),
612) -> Result<(BakeReport, BakeArtifacts), String> {
613 let t0 = std::time::Instant::now();
614 let o1_off = crate::nystrom::O1Cfg {
615 layers: crate::nystrom::O1Layers::List(Vec::new()),
616 m: 4,
617 w: 8,
618 sink: 1,
619 rect: crate::nystrom::O1_DEFAULT_RECT,
620 };
621 let fm = FcdModel::from_cmf(model, &o1_off, false)?;
622 let nl = fm.layers.len();
623 let inter = fm.layers.iter().map(|l| l.inter).max().unwrap_or(0);
624 if fm.layers.iter().any(|l| l.inter != inter) {
625 return Err("skill bake: non-uniform FFN widths".into());
626 }
627 let held: Vec<Vec<u32>> = chunks[..held_n.min(chunks.len())].to_vec();
628 let calib: Vec<Vec<u32>> = chunks[held_n.min(chunks.len())..].to_vec();
629 if calib.len() < 12 {
630 return Err(format!(
631 "skill bake: corpus too small ({} calib chunks)",
632 calib.len()
633 ));
634 }
635 let fcd: Vec<usize> = (nl.saturating_sub(hy.fcd_layers)..nl).collect();
636 let _rng = SplitMix64::new(hy.seed);
637
638 let loops = fm.loops.max(1);
666 let m0 = mask_init_logit(loops);
667 let vn = nl * loops;
669 let mut logits: Vec<Vec<f32>> = vec![vec![m0; inter]; vn];
670 let mut ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
671
672 let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
675 let base_pass = Pass {
676 fm: &fm,
677 tau: hy.tau,
678 logits: &open,
679 hard: true,
680 ffn: &ffn,
681 };
682 let backbone = held_ppl(&base_pass, &held);
683 log(&format!("baseline (full): {backbone:.3}"));
684
685 let mut adam_a = Adam::new(&vec![inter; vn], hy.lr_a);
687 let mut l1 = hy.l1_init * hy.l1_mult;
688 let l1_step_eff = hy.l1_step * hy.l1_mult;
689 let mut best: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
691 let mut max_sp: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
693 let mut prev_alive: Option<Vec<Vec<bool>>> = None;
694 let mut acc_chunk = 0f64;
698 let mut acc_adam = 0f64;
699 crate::gpu::bake_precision_strict(true);
707 for step in 0..hy.steps_a {
708 let t_step = std::time::Instant::now();
709 let chunk = &calib[step % calib.len()];
710 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
711 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = vec![None; nl];
712 let pass = Pass {
713 fm: &fm,
714 tau: hy.tau,
715 logits: &logits,
716 hard: false,
717 ffn: &ffn,
718 };
719 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
720 let l1_per = l1 / (inter as f64 * nl as f64);
722 for li in 0..vn {
723 for j in 0..inter {
724 let s = sigmoid(logits[li][j]) as f64;
725 dmask[li][j] = dmask[li][j] * s * (1.0 - s) + l1_per * s * (1.0 - s);
726 }
727 }
728 let t_chunk = t_step.elapsed().as_secs_f64();
741 let mut params: Vec<&mut [f32]> = logits.iter_mut().map(|v| v.as_mut_slice()).collect();
742 adam_a.step(&mut params, &dmask, 1.0);
743 acc_chunk += t_chunk;
744 acc_adam += t_step.elapsed().as_secs_f64() - t_chunk;
745 if (step + 1) % hy.eval_every == 0 {
746 l1 += l1_step_eff;
747 let pass = Pass {
748 fm: &fm,
749 tau: hy.tau,
750 logits: &logits,
751 hard: true,
752 ffn: &ffn,
753 };
754 crate::gpu::bake_precision_strict(false);
755 let hp = held_ppl(&pass, &held);
756 crate::gpu::bake_precision_strict(true);
757 let cur: Vec<Vec<bool>> = logits
763 .iter()
764 .map(|l| l.iter().map(|&x| sigmoid(x) > hy.tau).collect())
765 .collect();
766 if let Some(prev) = &prev_alive {
767 let died: Vec<String> = cur
768 .iter()
769 .zip(prev)
770 .enumerate()
771 .flat_map(|(li, (c, p))| {
772 c.iter()
773 .zip(p.iter())
774 .enumerate()
775 .filter(|&(_, (&cj, &pj))| pj && !cj)
776 .map(move |(j, _)| format!("L{li}:{j}"))
777 })
778 .collect();
779 if !died.is_empty() {
780 log(&format!(
781 " closed since last eval: {}: {}{}",
782 died.len(),
783 died.iter().take(32).cloned().collect::<Vec<_>>().join(" "),
784 if died.len() > 32 { " …" } else { "" }
785 ));
786 }
787 }
788 let alive: usize = cur.iter().map(|l| l.iter().filter(|&&b| b).count()).sum();
789 prev_alive = Some(cur);
790 let sp = 1.0 - alive as f64 / (vn * inter) as f64;
791 if sp > max_sp.2 {
793 max_sp = (hp, Some(logits.clone()), sp);
794 }
795 if hy.target_sparsity > 0.0 {
797 if sp >= hy.target_sparsity && hp < best.0 {
798 best = (hp, Some(logits.clone()), sp);
799 }
800 } else if hp < best.0 {
801 best = (hp, Some(logits.clone()), sp);
802 }
803 log(&format!(
804 " [A] step {}: L1={l1:.3} pruned={:.2}% hard-PPL={hp:.3} (bottom {}@{:.2}%) [fwd+bwd {:.1}s, adam {:.2}s per step]",
805 step + 1,
806 sp * 100.0,
807 if best.0 == f64::MAX {
808 "—".to_string()
809 } else {
810 format!("{:.3}", best.0)
811 },
812 best.2 * 100.0,
813 acc_chunk / (step + 1) as f64,
814 acc_adam / (step + 1) as f64
815 ));
816 }
817 }
818 crate::gpu::bake_precision_strict(false);
822 if hy.target_sparsity > 0.0 && best.1.is_none() {
823 log(&format!(
824 "[A] target sparsity {:.0}% not reached; using max-sparsity checkpoint ({:.0}%)",
825 hy.target_sparsity * 100.0,
826 max_sp.2 * 100.0
827 ));
828 best = max_sp;
829 }
830 {
833 use crate::fcd::prof;
834 let (a, f, bw, g, gc) = (
835 prof::take(&prof::ATTN_FWD),
836 prof::take(&prof::FFN_FWD),
837 prof::take(&prof::BWD),
838 prof::take(&prof::GEMM),
839 prof::GEMM_CALLS.swap(0, std::sync::atomic::Ordering::Relaxed),
840 );
841 log(&format!(
842 "[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)",
843 hy.steps_a,
844 if gc > 0 { g * 1000.0 / gc as f64 } else { 0.0 }
845 ));
846 log(&format!("[prof] gemm shapes:\n{}", prof::shape_report(6)));
847 }
848 if let Some(b) = best.1.take() {
849 logits = b;
850 }
851 let pass = Pass {
852 fm: &fm,
853 tau: hy.tau,
854 logits: &logits,
855 hard: true,
856 ffn: &ffn,
857 };
858 let masked = held_ppl(&pass, &held);
859 log(&format!(
860 "[A] {:.0}s: masked-PPL {masked:.3}",
861 t0.elapsed().as_secs_f64()
862 ));
863
864 for &li in &fcd {
866 let p = format!("model.layers.{li}.");
867 ffn[li] = Some((
868 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.gate_proj.weight"))
869 .map_err(|e| format!("phase-B gate: {e}"))?,
870 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.up_proj.weight"))
871 .map_err(|e| format!("phase-B up: {e}"))?,
872 crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.down_proj.weight"))
873 .map_err(|e| format!("phase-B down: {e}"))?,
874 ));
875 }
876 let sizes: Vec<usize> = fcd
877 .iter()
878 .flat_map(|&li| {
879 let (g, u, d) = ffn[li].as_ref().expect("phase-B masters");
880 [g.len(), u.len(), d.len()]
881 })
882 .collect();
883 let mut adam_b = Adam::new(&sizes, hy.lr_b);
884 let mut best_b: (f64, Option<Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>>>) = (masked, None);
885 for step in 0..hy.steps_b {
886 let chunk = &calib[step % calib.len()];
887 let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
888 let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = (0..nl)
889 .map(|li| {
890 ffn[li]
891 .as_ref()
892 .map(|(g, u, d)| (vec![0.0; g.len()], vec![0.0; u.len()], vec![0.0; d.len()]))
893 })
894 .collect();
895 let pass = Pass {
896 fm: &fm,
897 tau: hy.tau,
898 logits: &logits,
899 hard: true,
900 ffn: &ffn,
901 };
902 let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
903 let lr_scale = 0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
905 let first_fcd = fcd[0];
906 let mut params: Vec<&mut [f32]> = Vec::new();
907 let mut grads: Vec<Vec<f64>> = Vec::new();
908 for (off, slot) in ffn[first_fcd..].iter_mut().enumerate() {
909 let li = first_fcd + off;
910 let Some((g, u, d)) = slot.as_mut() else {
911 continue;
912 };
913 let (dg, du, dd) = dffn[li].take().unwrap();
914 params.push(g.as_mut_slice());
915 grads.push(dg);
916 params.push(u.as_mut_slice());
917 grads.push(du);
918 params.push(d.as_mut_slice());
919 grads.push(dd);
920 }
921 adam_b.step(&mut params, &grads, lr_scale * mask_step_scale(loops));
926 if (step + 1) % hy.eval_every == 0 {
927 let pass = Pass {
928 fm: &fm,
929 tau: hy.tau,
930 logits: &logits,
931 hard: true,
932 ffn: &ffn,
933 };
934 let cur = held_ppl(&pass, &held);
935 if cur < best_b.0 {
936 best_b = (cur, Some(ffn.clone()));
937 }
938 log(&format!(
939 " [B] step {}: held-PPL {cur:.3} (best {:.3})",
940 step + 1,
941 best_b.0
942 ));
943 }
944 }
945 if let Some(b) = best_b.1.take() {
946 ffn = b;
947 }
948 let overlaid = best_b.0;
949
950 let keep_visits = keep_masks(&logits, hy.tau, hy.align, hy.uniform_inter);
955 let keep: Vec<Vec<bool>> = (0..nl)
956 .map(|li| {
957 (0..inter)
958 .map(|j| (0..loops).any(|v| keep_visits[v * nl + li][j]))
959 .collect()
960 })
961 .collect();
962 if hy.align > 1 || hy.uniform_inter {
963 let raw: usize = logits
964 .iter()
965 .map(|l| l.iter().filter(|&&x| sigmoid(x) > hy.tau).count())
966 .sum();
967 let padded: usize = keep_visits
971 .iter()
972 .map(|a| a.iter().filter(|&&x| x).count())
973 .sum::<usize>()
974 .saturating_sub(raw);
975 log(&format!(
976 "align: +{padded} neurons resurrected (align {}, uniform {})",
977 hy.align, hy.uniform_inter
978 ));
979 }
980 let mut down_out = Vec::with_capacity(nl);
981 let mut gate_up = Vec::with_capacity(nl);
982 let mut kept_per_layer = Vec::with_capacity(nl);
983 for li in 0..nl {
984 let alive = &keep[li];
985 kept_per_layer.push(alive.iter().filter(|&&a| a).count());
986 let mut down = match &ffn[li] {
987 Some((_, _, d)) => d.clone(),
988 None => fm.mats(li).expect("layer mats").down.clone(),
989 };
990 let hsz = fm.hidden;
991 for r in 0..hsz {
992 for (c, &a) in alive.iter().enumerate() {
993 if !a {
994 down[r * inter + c] = 0.0;
995 }
996 }
997 }
998 gate_up.push(ffn[li].as_ref().map(|(g, u, _)| (g.clone(), u.clone())));
999 down_out.push(down);
1000 }
1001 let total: usize = keep_visits
1002 .iter()
1003 .map(|a| a.iter().filter(|&&x| x).count())
1004 .sum();
1005 let report = BakeReport {
1006 backbone,
1007 masked,
1008 overlaid,
1009 pruned_ratio: 1.0 - total as f64 / (vn * inter) as f64,
1010 kept_per_layer,
1011 sec: t0.elapsed().as_secs_f64(),
1012 };
1013 let arts = BakeArtifacts {
1014 keep,
1015 keep_visits,
1016 down: down_out,
1017 gate_up,
1018 fcd_layers: fcd,
1019 logits: logits.clone(),
1020 };
1021 Ok((report, arts))
1022}
1023
1024fn keep_masks(logits: &[Vec<f32>], tau: f32, align: usize, uniform: bool) -> Vec<Vec<bool>> {
1032 let inter = logits[0].len();
1033 let round = |n: usize| -> usize {
1034 let n = n.max(1);
1035 if align <= 1 {
1036 n.min(inter)
1037 } else {
1038 (n.div_ceil(align) * align).min(inter)
1039 }
1040 };
1041 let mut want: Vec<usize> = logits
1042 .iter()
1043 .map(|l| round(l.iter().filter(|&&x| sigmoid(x) > tau).count()))
1044 .collect();
1045 if uniform {
1046 let k = want.iter().copied().max().unwrap_or(inter);
1047 want = vec![k; logits.len()];
1048 }
1049 logits
1050 .iter()
1051 .zip(&want)
1052 .map(|(l, &k)| {
1053 let mut idx: Vec<usize> = (0..inter).collect();
1054 idx.sort_unstable_by(|&a, &b| l[b].total_cmp(&l[a]));
1055 let mut alive = vec![false; inter];
1056 for &i in idx.iter().take(k) {
1057 alive[i] = true;
1058 }
1059 alive
1060 })
1061 .collect()
1062}
1063
1064#[cfg(test)]
1065mod tests {
1066 use super::*;
1067
1068 fn kept(masks: &[Vec<bool>]) -> Vec<usize> {
1069 masks
1070 .iter()
1071 .map(|m| m.iter().filter(|&&a| a).count())
1072 .collect()
1073 }
1074
1075 #[test]
1078 fn keep_masks_aligns_up_and_preserves_alive() {
1079 let inter = 96;
1080 let l0: Vec<f32> = (0..inter)
1083 .map(|i| if i < 40 { 1.0 } else { -1.0 - i as f32 * 0.01 })
1084 .collect();
1085 let l1: Vec<f32> = (0..inter)
1087 .map(|i| if i < 64 { 2.0 } else { -3.0 })
1088 .collect();
1089 let masks = keep_masks(&[l0.clone(), l1], 0.5, 32, false);
1090 assert_eq!(kept(&masks), vec![64, 64]);
1091 for i in 0..64 {
1094 assert!(masks[0][i], "neuron {i} should be kept");
1095 }
1096 for i in 64..inter {
1097 assert!(!masks[0][i], "neuron {i} should stay pruned");
1098 }
1099 }
1100
1101 #[test]
1103 fn keep_masks_uniform_takes_max() {
1104 let inter = 96;
1105 let l0: Vec<f32> = (0..inter)
1106 .map(|i| if i < 10 { 1.0 } else { -2.0 })
1107 .collect();
1108 let l1: Vec<f32> = (0..inter)
1109 .map(|i| if i < 70 { 1.0 } else { -2.0 })
1110 .collect();
1111 let masks = keep_masks(&[l0, l1], 0.5, 32, true);
1112 assert_eq!(kept(&masks), vec![96, 96]);
1113 }
1114
1115 #[test]
1118 fn keep_masks_edges() {
1119 let inter = 48;
1120 let l: Vec<f32> = (0..inter)
1121 .map(|i| if i < 47 { 1.0 } else { -2.0 })
1122 .collect();
1123 let masks = keep_masks(&[l.clone()], 0.5, 32, false);
1124 assert_eq!(kept(&masks), vec![48]); let masks = keep_masks(&[l], 0.5, 1, false);
1126 assert_eq!(kept(&masks), vec![47]);
1127 let dead: Vec<f32> = vec![-5.0; inter];
1128 let masks = keep_masks(&[dead], 0.5, 32, false);
1129 assert_eq!(kept(&masks), vec![32]); }
1131}