use crate::fcd::{FcdModel, LnFfn};
use crate::fcd_ops as ops;
use crate::sampler::SplitMix64;
use cortiq_core::CmfModel;
use std::sync::Arc;
#[derive(Clone, Debug)]
pub struct BakeHyper {
pub steps_a: usize,
pub steps_b: usize,
pub l1_init: f64,
pub l1_step: f64,
pub eval_every: usize,
pub lr_a: f64,
pub lr_b: f64,
pub tau: f32,
pub fcd_layers: usize,
pub seed: u64,
pub target_sparsity: f64,
pub l1_mult: f64,
pub align: usize,
pub uniform_inter: bool,
}
impl Default for BakeHyper {
fn default() -> Self {
Self {
steps_a: 240,
steps_b: 120,
l1_init: 0.01,
l1_step: 0.005,
eval_every: 30,
lr_a: 0.1,
lr_b: 1e-5,
tau: 0.5,
fcd_layers: 4,
seed: 0,
target_sparsity: 0.0,
l1_mult: 1.0,
align: 32,
uniform_inter: false,
}
}
}
pub struct BakeReport {
pub backbone: f64,
pub masked: f64,
pub overlaid: f64,
pub pruned_ratio: f64,
pub kept_per_layer: Vec<usize>,
pub sec: f64,
}
pub struct BakeArtifacts {
pub keep: Vec<Vec<bool>>,
pub keep_visits: Vec<Vec<bool>>,
pub down: Vec<Vec<f32>>,
pub gate_up: Vec<Option<(Vec<f32>, Vec<f32>)>>,
pub fcd_layers: Vec<usize>,
}
const CLIP: f64 = 1.0;
const B1: f64 = 0.9;
const B2: f64 = 0.999;
const EPS: f64 = 1e-8;
struct Adam {
m: Vec<Vec<f64>>,
v: Vec<Vec<f64>>,
t: i32,
lr: f64,
}
impl Adam {
fn new(sizes: &[usize], lr: f64) -> Self {
Self {
m: sizes.iter().map(|&n| vec![0.0; n]).collect(),
v: sizes.iter().map(|&n| vec![0.0; n]).collect(),
t: 0,
lr,
}
}
fn step(&mut self, params: &mut [&mut [f32]], grads: &[Vec<f64>], lr_scale: f64) {
let gn: f64 = grads
.iter()
.flat_map(|g| g.iter().map(|x| x * x))
.sum::<f64>()
.sqrt();
let clip = if gn > CLIP { CLIP / gn } else { 1.0 };
self.t += 1;
let (bc1, bc2) = (1.0 - B1.powi(self.t), 1.0 - B2.powi(self.t));
for (pi, p) in params.iter_mut().enumerate() {
for j in 0..p.len() {
let g = grads[pi][j] * clip;
let m = &mut self.m[pi][j];
let v = &mut self.v[pi][j];
*m = B1 * *m + (1.0 - B1) * g;
*v = B2 * *v + (1.0 - B2) * g * g;
let upd = (*m / bc1) / ((*v / bc2).sqrt() + EPS);
p[j] -= (self.lr * lr_scale * upd) as f32;
}
}
}
}
pub fn mask_init_logit(loops: usize) -> f32 {
let base = 1.0f32 / (1.0 + (-2.0f32).exp());
let per_visit = base.powf(1.0 / loops.max(1) as f32);
(per_visit / (1.0 - per_visit)).ln()
}
pub fn mask_step_scale(loops: usize) -> f64 {
1.0 / loops.max(1) as f64
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
struct Pass<'a> {
fm: &'a FcdModel,
tau: f32,
logits: &'a [Vec<f32>],
hard: bool,
ffn: &'a [Option<(Vec<f32>, Vec<f32>, Vec<f32>)>],
}
impl Pass<'_> {
fn gates(&self, li: usize) -> Vec<f32> {
self.logits[li]
.iter()
.map(|&l| {
let s = sigmoid(l);
if self.hard {
if s > self.tau { 1.0 } else { 0.0 }
} else {
s
}
})
.collect()
}
fn wts<'b>(&'b self, li: usize, mats: &'b crate::fcd::LayerMats) -> LnFfn<'b> {
let l = &self.fm.layers[li];
match &self.ffn[li] {
Some((g, u, d)) => LnFfn {
iln: &l.iln,
pln: &l.pln,
gate: g,
up: u,
down: d,
gu: None,
},
None => LnFfn {
iln: &l.iln,
pln: &l.pln,
gate: &[],
up: &[],
down: &mats.down,
gu: Some(&mats.gu),
},
}
}
#[allow(clippy::too_many_arguments)]
fn chunk(
&self,
ids: &[u32],
grad: Option<(
&mut [Vec<f64>],
&mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
)>,
) -> (f64, usize) {
self.chunk_batch(ids, 1, grad)
}
fn chunk_batch(
&self,
ids: &[u32],
b: usize,
grad: Option<(
&mut [Vec<f64>],
&mut [Option<(Vec<f64>, Vec<f64>, Vec<f64>)>],
)>,
) -> (f64, usize) {
let fm = self.fm;
let hsz = fm.hidden;
debug_assert!(ids.len() % b.max(1) == 0, "ragged batch");
debug_assert!(grad.is_none() || b == 1, "grads are per-chunk");
let t = ids.len() / b.max(1);
let n = b * t;
let nl = fm.layers.len();
let mut h = vec![0f32; n * hsz];
for (r, &id) in ids.iter().enumerate() {
h[r * hsz..(r + 1) * hsz]
.copy_from_slice(&fm.embed[id as usize * hsz..(id as usize + 1) * hsz]);
}
let loops = fm.loops.max(1);
let vn = nl * loops;
let mut h_ins = Vec::with_capacity(vn);
let mut acts = Vec::with_capacity(vn);
let mut masks = Vec::with_capacity(vn);
let mut lnorms: Vec<Option<(Vec<f32>, Vec<f32>)>> = vec![None; vn];
for vl in 0..vn {
let li = vl % nl;
let g = self.gates(vl);
let mats_hold = fm.mats(li).expect("layer mats");
let wts = self.wts(li, &mats_hold);
let want = grad.is_some();
let (h2, a) = fm.layer_forward_scaled(li, &h, b, t, &wts, false, want, Some(&g));
h_ins.push(if want { h } else { Vec::new() });
acts.push(a);
masks.push(g);
h = h2;
if fm.loop_norm && li + 1 == nl && vl + 1 < vn {
let mut hn = vec![0f32; n * hsz];
let mut inv = vec![0f32; n];
ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
if want {
lnorms[vl] = Some((h, inv));
}
h = hn;
}
}
let mut hn = vec![0f32; n * hsz];
let mut inv = vec![0f32; n];
ops::rmsnorm_fwd(&h, &fm.final_norm, fm.eps, fm.gemma, &mut hn, &mut inv);
let lm: &[f32] = fm.lm_head.as_deref().unwrap_or(&fm.embed);
let vocab = lm.len() / hsz;
let pool = fm.pool.as_deref();
let mut nll = 0f64;
let mut dh_n = vec![0f32; n * hsz]; const POS_CHUNK: usize = 64;
let scored = b * (t - 1);
for bi in 0..b {
let base = bi * t;
let mut p0 = 0usize;
while p0 < t - 1 {
let pc = POS_CHUNK.min(t - 1 - p0);
let mut logits = vec![0f32; pc * vocab];
ops::gemm_nt(
&hn[(base + p0) * hsz..(base + p0 + pc) * hsz],
lm,
&mut logits,
pc,
hsz,
vocab,
pool,
);
for r in 0..pc {
let target = ids[base + p0 + r + 1] as usize;
let row = &mut logits[r * vocab..(r + 1) * vocab];
let mx = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max) as f64;
let mut sum = 0f64;
for v in row.iter() {
sum += ((*v as f64) - mx).exp();
}
nll += mx + sum.ln() - row[target] as f64;
if grad.is_some() {
let inv_n = 1.0 / scored as f64;
for v in row.iter_mut() {
*v = ((((*v as f64) - mx).exp() / sum) * inv_n) as f32;
}
row[target] -= inv_n as f32;
}
}
if grad.is_some() {
ops::gemm_dx(
&logits,
lm,
&mut dh_n[(base + p0) * hsz..(base + p0 + pc) * hsz],
pc,
hsz,
vocab,
pool,
);
}
p0 += pc;
}
}
let Some((dmask, dffn)) = grad else {
return (nll, scored);
};
let t_bwd = std::time::Instant::now();
let mut dh = vec![0f32; n * hsz];
ops::rmsnorm_bwd(&h, &fm.final_norm, &inv, &dh_n, fm.gemma, &mut dh, None);
for vl in (0..vn).rev() {
let li = vl % nl;
if let Some((hb, inv)) = lnorms[vl].as_ref() {
let mut dprev = vec![0f32; n * hsz];
ops::rmsnorm_bwd(hb, &fm.final_norm, inv, &dh, fm.gemma, &mut dprev, None);
dh = dprev;
}
let a = acts[vl].as_ref().expect("acts saved in grad mode");
let g = &masks[vl];
let inter = fm.layers[li].inter;
let mats_hold = fm.mats(li).expect("layer mats");
let wts = self.wts(li, &mats_hold);
let mut dact2 = vec![0f32; t * inter];
ops::gemm_dx(&dh, wts.down, &mut dact2, t, inter, hsz, fm.pool.as_deref());
if let Some((_, _, dd)) = dffn[li].as_mut() {
let mut act2 = a.act.clone();
for r in 0..t {
for (x, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
*x *= gv;
}
}
let mut dw = vec![0f32; hsz * inter];
ops::gemm_dw(&dh, &act2, &mut dw, t, inter, hsz, fm.pool.as_deref());
for (o, &x) in dd.iter_mut().zip(&dw) {
*o += x as f64;
}
}
{
let dm = &mut dmask[vl];
for r in 0..t {
let da = &dact2[r * inter..(r + 1) * inter];
let aa = &a.act[r * inter..(r + 1) * inter];
for j in 0..inter {
dm[j] += da[j] as f64 * aa[j] as f64;
}
}
for (j, d) in dm.iter_mut().enumerate() {
let _ = j;
let _ = d;
}
}
let mut dg_pre = vec![0f32; t * inter];
let mut du_pre = vec![0f32; t * inter];
for r in 0..t {
for j in 0..inter {
let i = r * inter + j;
let da = dact2[i] * g[j];
let sg = ops::silu(a.gpre[i]);
dg_pre[i] = da * a.upre[i] * ops::silu_bwd(a.gpre[i]);
du_pre[i] = da * sg;
}
}
let mut dn2 = vec![0f32; t * hsz];
if let Some(gu) = wts.gu {
let mut dgu = vec![0f32; t * 2 * inter];
for r in 0..t {
let row = &mut dgu[r * 2 * inter..(r + 1) * 2 * inter];
row[..inter].copy_from_slice(&dg_pre[r * inter..(r + 1) * inter]);
row[inter..].copy_from_slice(&du_pre[r * inter..(r + 1) * inter]);
}
ops::gemm_dx(&dgu, gu, &mut dn2, t, hsz, 2 * inter, fm.pool.as_deref());
} else {
ops::gemm_dx(
&dg_pre,
wts.gate,
&mut dn2,
t,
hsz,
inter,
fm.pool.as_deref(),
);
let mut dn2b = vec![0f32; t * hsz];
ops::gemm_dx(
&du_pre,
wts.up,
&mut dn2b,
t,
hsz,
inter,
fm.pool.as_deref(),
);
for (x, &y) in dn2.iter_mut().zip(&dn2b) {
*x += y;
}
}
if let Some((dgw, duw, _)) = dffn[li].as_mut() {
let mut dw = vec![0f32; inter * hsz];
ops::gemm_dw(&dg_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
for (o, &x) in dgw.iter_mut().zip(&dw) {
*o += x as f64;
}
dw.fill(0.0);
ops::gemm_dw(&du_pre, &a.n2, &mut dw, t, hsz, inter, fm.pool.as_deref());
for (o, &x) in duw.iter_mut().zip(&dw) {
*o += x as f64;
}
}
let mut dh1 = dh.clone(); ops::rmsnorm_bwd(&a.h1, wts.pln, &a.inv2, &dn2, fm.gemma, &mut dh1, None);
dh = dh1;
let _ = &h_ins[vl];
}
crate::fcd::prof::add(&crate::fcd::prof::BWD, t_bwd);
(nll, scored)
}
}
fn held_ppl(pass: &Pass, held: &[Vec<u32>]) -> f64 {
if held.is_empty() {
return f64::NAN;
}
let t = held[0].len();
if held.iter().all(|c| c.len() == t) {
let flat: Vec<u32> = held.iter().flatten().copied().collect();
let (l, k) = pass.chunk_batch(&flat, held.len(), None);
return (l / k.max(1) as f64).exp();
}
let mut nll = 0f64;
let mut n = 0usize;
for c in held {
let (l, k) = pass.chunk(c, None);
nll += l;
n += k;
}
(nll / n.max(1) as f64).exp()
}
pub fn skill_bake(
model: &Arc<CmfModel>,
chunks: &[Vec<u32>],
held_n: usize,
hy: &BakeHyper,
mut log: impl FnMut(&str),
) -> Result<(BakeReport, BakeArtifacts), String> {
let t0 = std::time::Instant::now();
let o1_off = crate::nystrom::O1Cfg {
layers: crate::nystrom::O1Layers::List(Vec::new()),
m: 4,
w: 8,
sink: 1,
rect: crate::nystrom::O1_DEFAULT_RECT,
};
let fm = FcdModel::from_cmf(model, &o1_off)?;
let nl = fm.layers.len();
let inter = fm.layers.iter().map(|l| l.inter).max().unwrap_or(0);
if fm.layers.iter().any(|l| l.inter != inter) {
return Err("skill bake: non-uniform FFN widths".into());
}
let held: Vec<Vec<u32>> = chunks[..held_n.min(chunks.len())].to_vec();
let calib: Vec<Vec<u32>> = chunks[held_n.min(chunks.len())..].to_vec();
if calib.len() < 12 {
return Err(format!(
"skill bake: corpus too small ({} calib chunks)",
calib.len()
));
}
let fcd: Vec<usize> = (nl.saturating_sub(hy.fcd_layers)..nl).collect();
let _rng = SplitMix64::new(hy.seed);
let loops = fm.loops.max(1);
let m0 = mask_init_logit(loops);
let vn = nl * loops;
let mut logits: Vec<Vec<f32>> = vec![vec![m0; inter]; vn];
let mut ffn: Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>> = vec![None; nl];
let open: Vec<Vec<f32>> = vec![vec![50.0; inter]; vn];
let base_pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &open,
hard: true,
ffn: &ffn,
};
let backbone = held_ppl(&base_pass, &held);
log(&format!("baseline (full): {backbone:.3}"));
let mut adam_a = Adam::new(&vec![inter; vn], hy.lr_a);
let mut l1 = hy.l1_init * hy.l1_mult;
let l1_step_eff = hy.l1_step * hy.l1_mult;
let mut best: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
let mut max_sp: (f64, Option<Vec<Vec<f32>>>, f64) = (f64::MAX, None, 0.0);
let mut prev_alive: Option<Vec<Vec<bool>>> = None;
let mut acc_chunk = 0f64;
let mut acc_adam = 0f64;
crate::gpu::bake_precision_strict(true);
for step in 0..hy.steps_a {
let t_step = std::time::Instant::now();
let chunk = &calib[step % calib.len()];
let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = vec![None; nl];
let pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &logits,
hard: false,
ffn: &ffn,
};
let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
let l1_per = l1 / (inter as f64 * nl as f64);
for li in 0..vn {
for j in 0..inter {
let s = sigmoid(logits[li][j]) as f64;
dmask[li][j] = dmask[li][j] * s * (1.0 - s) + l1_per * s * (1.0 - s);
}
}
let t_chunk = t_step.elapsed().as_secs_f64();
let mut params: Vec<&mut [f32]> = logits.iter_mut().map(|v| v.as_mut_slice()).collect();
adam_a.step(&mut params, &dmask, 1.0);
acc_chunk += t_chunk;
acc_adam += t_step.elapsed().as_secs_f64() - t_chunk;
if (step + 1) % hy.eval_every == 0 {
l1 += l1_step_eff;
let pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &logits,
hard: true,
ffn: &ffn,
};
crate::gpu::bake_precision_strict(false);
let hp = held_ppl(&pass, &held);
crate::gpu::bake_precision_strict(true);
let cur: Vec<Vec<bool>> = logits
.iter()
.map(|l| l.iter().map(|&x| sigmoid(x) > hy.tau).collect())
.collect();
if let Some(prev) = &prev_alive {
let died: Vec<String> = cur
.iter()
.zip(prev)
.enumerate()
.flat_map(|(li, (c, p))| {
c.iter()
.zip(p.iter())
.enumerate()
.filter(|&(_, (&cj, &pj))| pj && !cj)
.map(move |(j, _)| format!("L{li}:{j}"))
})
.collect();
if !died.is_empty() {
log(&format!(
" closed since last eval: {}: {}{}",
died.len(),
died.iter().take(32).cloned().collect::<Vec<_>>().join(" "),
if died.len() > 32 { " …" } else { "" }
));
}
}
let alive: usize = cur.iter().map(|l| l.iter().filter(|&&b| b).count()).sum();
prev_alive = Some(cur);
let sp = 1.0 - alive as f64 / (vn * inter) as f64;
if sp > max_sp.2 {
max_sp = (hp, Some(logits.clone()), sp);
}
if hy.target_sparsity > 0.0 {
if sp >= hy.target_sparsity && hp < best.0 {
best = (hp, Some(logits.clone()), sp);
}
} else if hp < best.0 {
best = (hp, Some(logits.clone()), sp);
}
log(&format!(
" [A] step {}: L1={l1:.3} pruned={:.2}% hard-PPL={hp:.3} (bottom {}@{:.2}%) [fwd+bwd {:.1}s, adam {:.2}s per step]",
step + 1,
sp * 100.0,
if best.0 == f64::MAX {
"—".to_string()
} else {
format!("{:.3}", best.0)
},
best.2 * 100.0,
acc_chunk / (step + 1) as f64,
acc_adam / (step + 1) as f64
));
}
}
crate::gpu::bake_precision_strict(false);
if hy.target_sparsity > 0.0 && best.1.is_none() {
log(&format!(
"[A] target sparsity {:.0}% not reached; using max-sparsity checkpoint ({:.0}%)",
hy.target_sparsity * 100.0,
max_sp.2 * 100.0
));
best = max_sp;
}
{
use crate::fcd::prof;
let (a, f, bw, g, gc) = (
prof::take(&prof::ATTN_FWD),
prof::take(&prof::FFN_FWD),
prof::take(&prof::BWD),
prof::take(&prof::GEMM),
prof::GEMM_CALLS.swap(0, std::sync::atomic::Ordering::Relaxed),
);
log(&format!(
"[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)",
hy.steps_a,
if gc > 0 { g * 1000.0 / gc as f64 } else { 0.0 }
));
log(&format!("[prof] gemm shapes:\n{}", prof::shape_report(6)));
}
if let Some(b) = best.1.take() {
logits = b;
}
let pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &logits,
hard: true,
ffn: &ffn,
};
let masked = held_ppl(&pass, &held);
log(&format!(
"[A] {:.0}s: masked-PPL {masked:.3}",
t0.elapsed().as_secs_f64()
));
for &li in &fcd {
let p = format!("model.layers.{li}.");
ffn[li] = Some((
crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.gate_proj.weight"))
.map_err(|e| format!("phase-B gate: {e}"))?,
crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.up_proj.weight"))
.map_err(|e| format!("phase-B up: {e}"))?,
crate::fcd::deq_pub(&fm.src, &format!("{p}mlp.down_proj.weight"))
.map_err(|e| format!("phase-B down: {e}"))?,
));
}
let sizes: Vec<usize> = fcd
.iter()
.flat_map(|&li| {
let (g, u, d) = ffn[li].as_ref().expect("phase-B masters");
[g.len(), u.len(), d.len()]
})
.collect();
let mut adam_b = Adam::new(&sizes, hy.lr_b);
let mut best_b: (f64, Option<Vec<Option<(Vec<f32>, Vec<f32>, Vec<f32>)>>>) = (masked, None);
for step in 0..hy.steps_b {
let chunk = &calib[step % calib.len()];
let mut dmask: Vec<Vec<f64>> = vec![vec![0.0; inter]; vn];
let mut dffn: Vec<Option<(Vec<f64>, Vec<f64>, Vec<f64>)>> = (0..nl)
.map(|li| {
ffn[li]
.as_ref()
.map(|(g, u, d)| (vec![0.0; g.len()], vec![0.0; u.len()], vec![0.0; d.len()]))
})
.collect();
let pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &logits,
hard: true,
ffn: &ffn,
};
let _ = pass.chunk(chunk, Some((&mut dmask, &mut dffn)));
let lr_scale = 0.5 * (1.0 + (std::f64::consts::PI * step as f64 / hy.steps_b as f64).cos());
let first_fcd = fcd[0];
let mut params: Vec<&mut [f32]> = Vec::new();
let mut grads: Vec<Vec<f64>> = Vec::new();
for (off, slot) in ffn[first_fcd..].iter_mut().enumerate() {
let li = first_fcd + off;
let Some((g, u, d)) = slot.as_mut() else {
continue;
};
let (dg, du, dd) = dffn[li].take().unwrap();
params.push(g.as_mut_slice());
grads.push(dg);
params.push(u.as_mut_slice());
grads.push(du);
params.push(d.as_mut_slice());
grads.push(dd);
}
adam_b.step(&mut params, &grads, lr_scale * mask_step_scale(loops));
if (step + 1) % hy.eval_every == 0 {
let pass = Pass {
fm: &fm,
tau: hy.tau,
logits: &logits,
hard: true,
ffn: &ffn,
};
let cur = held_ppl(&pass, &held);
if cur < best_b.0 {
best_b = (cur, Some(ffn.clone()));
}
log(&format!(
" [B] step {}: held-PPL {cur:.3} (best {:.3})",
step + 1,
best_b.0
));
}
}
if let Some(b) = best_b.1.take() {
ffn = b;
}
let overlaid = best_b.0;
let keep_visits = keep_masks(&logits, hy.tau, hy.align, hy.uniform_inter);
let keep: Vec<Vec<bool>> = (0..nl)
.map(|li| {
(0..inter)
.map(|j| (0..loops).any(|v| keep_visits[v * nl + li][j]))
.collect()
})
.collect();
if hy.align > 1 || hy.uniform_inter {
let raw: usize = logits
.iter()
.map(|l| l.iter().filter(|&&x| sigmoid(x) > hy.tau).count())
.sum();
let padded: usize = keep_visits
.iter()
.map(|a| a.iter().filter(|&&x| x).count())
.sum::<usize>()
.saturating_sub(raw);
log(&format!(
"align: +{padded} neurons resurrected (align {}, uniform {})",
hy.align, hy.uniform_inter
));
}
let mut down_out = Vec::with_capacity(nl);
let mut gate_up = Vec::with_capacity(nl);
let mut kept_per_layer = Vec::with_capacity(nl);
for li in 0..nl {
let alive = &keep[li];
kept_per_layer.push(alive.iter().filter(|&&a| a).count());
let mut down = match &ffn[li] {
Some((_, _, d)) => d.clone(),
None => fm.mats(li).expect("layer mats").down.clone(),
};
let hsz = fm.hidden;
for r in 0..hsz {
for (c, &a) in alive.iter().enumerate() {
if !a {
down[r * inter + c] = 0.0;
}
}
}
gate_up.push(ffn[li].as_ref().map(|(g, u, _)| (g.clone(), u.clone())));
down_out.push(down);
}
let total: usize = keep_visits
.iter()
.map(|a| a.iter().filter(|&&x| x).count())
.sum();
let report = BakeReport {
backbone,
masked,
overlaid,
pruned_ratio: 1.0 - total as f64 / (vn * inter) as f64,
kept_per_layer,
sec: t0.elapsed().as_secs_f64(),
};
let arts = BakeArtifacts {
keep,
keep_visits,
down: down_out,
gate_up,
fcd_layers: fcd,
};
Ok((report, arts))
}
fn keep_masks(logits: &[Vec<f32>], tau: f32, align: usize, uniform: bool) -> Vec<Vec<bool>> {
let inter = logits[0].len();
let round = |n: usize| -> usize {
let n = n.max(1);
if align <= 1 {
n.min(inter)
} else {
(n.div_ceil(align) * align).min(inter)
}
};
let mut want: Vec<usize> = logits
.iter()
.map(|l| round(l.iter().filter(|&&x| sigmoid(x) > tau).count()))
.collect();
if uniform {
let k = want.iter().copied().max().unwrap_or(inter);
want = vec![k; logits.len()];
}
logits
.iter()
.zip(&want)
.map(|(l, &k)| {
let mut idx: Vec<usize> = (0..inter).collect();
idx.sort_unstable_by(|&a, &b| l[b].total_cmp(&l[a]));
let mut alive = vec![false; inter];
for &i in idx.iter().take(k) {
alive[i] = true;
}
alive
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn kept(masks: &[Vec<bool>]) -> Vec<usize> {
masks
.iter()
.map(|m| m.iter().filter(|&&a| a).count())
.collect()
}
#[test]
fn keep_masks_aligns_up_and_preserves_alive() {
let inter = 96;
let l0: Vec<f32> = (0..inter)
.map(|i| if i < 40 { 1.0 } else { -1.0 - i as f32 * 0.01 })
.collect();
let l1: Vec<f32> = (0..inter)
.map(|i| if i < 64 { 2.0 } else { -3.0 })
.collect();
let masks = keep_masks(&[l0.clone(), l1], 0.5, 32, false);
assert_eq!(kept(&masks), vec![64, 64]);
for i in 0..64 {
assert!(masks[0][i], "neuron {i} should be kept");
}
for i in 64..inter {
assert!(!masks[0][i], "neuron {i} should stay pruned");
}
}
#[test]
fn keep_masks_uniform_takes_max() {
let inter = 96;
let l0: Vec<f32> = (0..inter)
.map(|i| if i < 10 { 1.0 } else { -2.0 })
.collect();
let l1: Vec<f32> = (0..inter)
.map(|i| if i < 70 { 1.0 } else { -2.0 })
.collect();
let masks = keep_masks(&[l0, l1], 0.5, 32, true);
assert_eq!(kept(&masks), vec![96, 96]);
}
#[test]
fn keep_masks_edges() {
let inter = 48;
let l: Vec<f32> = (0..inter)
.map(|i| if i < 47 { 1.0 } else { -2.0 })
.collect();
let masks = keep_masks(&[l.clone()], 0.5, 32, false);
assert_eq!(kept(&masks), vec![48]); let masks = keep_masks(&[l], 0.5, 1, false);
assert_eq!(kept(&masks), vec![47]);
let dead: Vec<f32> = vec![-5.0; inter];
let masks = keep_masks(&[dead], 0.5, 32, false);
assert_eq!(kept(&masks), vec![32]); }
}