1use crate::fcd_ops::{self as ops, NysCfg};
25use crate::nystrom::{O1Cfg, O1Layers};
26use crate::pipeline::{DenseFfn, FfnKind, Pipeline};
27use crate::pool::Pool;
28use crate::qtensor::QTensor;
29use crate::sampler::{SamplerConfig, SplitMix64};
30use cortiq_core::{CmfModel, LayerType, NormStyle, TensorDtype};
31use std::sync::Arc;
32
33const LM_CHUNK: usize = 32;
36
37const ADAM_B1: f64 = 0.9;
39const ADAM_B2: f64 = 0.999;
40const ADAM_EPS: f64 = 1e-8;
41const ADAM_WD: f64 = 0.01;
42
43#[derive(Clone, Debug)]
45pub struct FcdHyper {
46 pub steps: usize,
47 pub lr: f64,
48 pub kl_w: f64,
49 pub eval_every: usize,
50 pub bs: usize,
51 pub seq: usize,
52 pub seed: u64,
53}
54
55impl Default for FcdHyper {
56 fn default() -> Self {
57 Self {
58 steps: 300,
59 lr: 5e-5,
60 kl_w: 0.7,
61 eval_every: 25,
62 bs: 2,
63 seq: 512,
64 seed: 0,
65 }
66 }
67}
68
69#[derive(Clone, Debug)]
72pub struct FcdReport {
73 pub converted: Vec<usize>,
74 pub teacher_ppl: f64,
76 pub ppl_start: f64,
78 pub ppl_best: f64,
80 pub best_step: usize,
81 pub ppl_final: f64,
83 pub steps_run: usize,
84 pub sec_per_step: f64,
85 pub losses: Vec<(f64, f64)>,
87 pub gate: Option<GateReport>,
89}
90
91#[derive(Clone, Debug)]
93pub struct GateReport {
94 pub baseline: Vec<f64>,
96 pub evals: Vec<(usize, f64, Vec<f64>, bool)>,
98 pub chosen: Option<usize>,
101}
102
103pub fn loop_score(ids: &[u32]) -> f64 {
108 if ids.len() < 5 {
109 return 0.0;
110 }
111 let grams: std::collections::HashSet<&[u32]> = ids.windows(4).collect();
112 1.0 - grams.len() as f64 / ids.windows(4).count() as f64
113}
114
115#[derive(Clone, Debug)]
119pub struct GenGateCfg {
120 pub prompts: Vec<Vec<u32>>,
123 pub gen_tokens: usize,
124 pub threshold: f64,
126 pub baseline_slack: f64,
128}
129
130impl GenGateCfg {
131 pub fn standard(va: &[u32]) -> Option<Self> {
134 let l = va.len().saturating_sub(500);
135 if l < 400 {
136 return None;
137 }
138 let prompts = [l / 10, l / 2, 8 * l / 10]
139 .iter()
140 .map(|&off| va[off..off + 400].to_vec())
141 .collect();
142 Some(Self {
143 prompts,
144 gen_tokens: 60,
145 threshold: 0.35,
146 baseline_slack: 0.10,
147 })
148 }
149}
150
151pub fn gate_pass(scores: &[f64], baseline: &[f64], threshold: f64, slack: f64) -> bool {
155 scores
156 .iter()
157 .zip(baseline)
158 .all(|(&s, &b)| s <= threshold && s <= b + slack)
159}
160
161pub fn select_checkpoint(
166 evals: &[(usize, f64, Vec<f64>)],
167 baseline: &[f64],
168 threshold: f64,
169 slack: f64,
170) -> Option<usize> {
171 let mut best: Option<usize> = None;
172 for (i, (_, ppl, scores)) in evals.iter().enumerate() {
173 if !gate_pass(scores, baseline, threshold, slack) {
174 continue;
175 }
176 if best.map(|b| *ppl < evals[b].1).unwrap_or(true) {
177 best = Some(i);
178 }
179 }
180 best
181}
182
183pub mod prof {
189 use std::sync::atomic::{AtomicU64, Ordering};
190 pub static ATTN_FWD: AtomicU64 = AtomicU64::new(0);
191 pub static FFN_FWD: AtomicU64 = AtomicU64::new(0);
192 pub static BWD: AtomicU64 = AtomicU64::new(0);
193 pub static GEMM: AtomicU64 = AtomicU64::new(0);
194 pub static GEMM_CALLS: AtomicU64 = AtomicU64::new(0);
195 #[inline]
196 pub fn add(c: &AtomicU64, t: std::time::Instant) {
197 c.fetch_add(t.elapsed().as_nanos() as u64, Ordering::Relaxed);
198 }
199 pub fn take(c: &AtomicU64) -> f64 {
201 c.swap(0, Ordering::Relaxed) as f64 / 1e9
202 }
203
204 use std::sync::Mutex;
205 pub static SHAPES: Mutex<Vec<((usize, usize, usize), (u64, u64))>> = Mutex::new(Vec::new());
209 pub fn gemm_shape(n: usize, k: usize, m: usize, t0: std::time::Instant) {
210 let ns = t0.elapsed().as_nanos() as u64;
211 let mut g = SHAPES.lock().unwrap();
212 match g.iter_mut().find(|(s, _)| *s == (n, k, m)) {
213 Some((_, (c, tt))) => {
214 *c += 1;
215 *tt += ns;
216 }
217 None => g.push(((n, k, m), (1, ns))),
218 }
219 }
220 pub fn shape_report(top: usize) -> String {
221 let mut g = SHAPES.lock().unwrap();
222 g.sort_by_key(|(_, (_, ns))| std::cmp::Reverse(*ns));
223 let out = g
224 .iter()
225 .take(top)
226 .map(|((n, k, m), (c, ns))| {
227 format!(
228 " [{n}x{k}x{m}] {c} calls, {:.1}s total, {:.1} ms/call",
229 *ns as f64 / 1e9,
230 *ns as f64 / 1e6 / *c as f64
231 )
232 })
233 .collect::<Vec<_>>()
234 .join("\n");
235 g.clear();
236 out
237 }
238}
239
240enum FcdAttn {
243 Full {
244 wq: Vec<f32>,
245 wk: Vec<f32>,
246 wv: Vec<f32>,
247 wqkv: Vec<f32>,
251 wo: Vec<f32>,
252 q_norm: Option<Vec<f32>>,
253 k_norm: Option<Vec<f32>>,
254 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
255 output_gate: bool,
258 },
259 Gdn {
262 wqkv: Vec<f32>,
263 wz: Vec<f32>,
264 wa: Vec<f32>,
265 wb: Vec<f32>,
266 conv: Vec<f32>,
267 a_log: Vec<f32>,
268 dt_bias: Vec<f32>,
269 norm: Vec<f32>,
270 wout: Vec<f32>,
271 },
272}
273
274pub(crate) struct FcdLayer {
275 attn: FcdAttn,
276 pub(crate) inter: usize,
277 pub(crate) iln: Vec<f32>,
279 pub(crate) pln: Vec<f32>,
280 pub(crate) gate: Vec<f32>,
281 pub(crate) up: Vec<f32>,
282 pub(crate) down: Vec<f32>,
283 pub(crate) gu: Vec<f32>,
285}
286
287#[derive(Clone, Copy)]
289struct GdnDims {
290 nv: usize,
291 nk: usize,
292 dk: usize,
293 dv: usize,
294 kk: usize,
295}
296
297impl GdnDims {
298 fn c_dim(&self) -> usize {
299 2 * self.nk * self.dk + self.nv * self.dv
300 }
301 fn vd(&self) -> usize {
302 self.nv * self.dv
303 }
304}
305
306pub struct FcdModel {
308 pub hidden: usize,
309 pub nh: usize,
310 pub nkv: usize,
311 pub hd: usize,
312 pub nl: usize,
313 pub vocab: usize,
314 pub(crate) eps: f64,
315 pub(crate) gemma: bool,
316 rotary_dim: usize,
317 inv_freq: Vec<f64>,
318 pub(crate) embed: Vec<f32>,
320 pub(crate) lm_head: Option<Vec<f32>>,
321 pub(crate) final_norm: Vec<f32>,
322 pub(crate) layers: Vec<FcdLayer>,
323 o1_flags: Vec<bool>,
325 nys: NysCfg,
326 gdn: Option<GdnDims>,
328 pub(crate) loops: usize,
335 pub(crate) loop_norm: bool,
336 pub(crate) pool: Option<Arc<Pool>>,
337}
338
339fn deq(model: &CmfModel, name: &str) -> Result<Vec<f32>, String> {
340 let e = model
341 .tensor(name)
342 .ok_or_else(|| format!("tensor '{name}' not found"))?;
343 let mut out = vec![0f32; e.n_elems()];
344 cortiq_core::quant::dequant_tensor(e, model.entry_bytes(e), &mut out)?;
345 Ok(out)
346}
347
348impl FcdModel {
349 pub fn from_cmf(model: &CmfModel, o1: &O1Cfg) -> Result<Self, String> {
352 let arch = model.arch().clone();
353 if arch.hidden_act != "silu" {
354 return Err(format!(
355 "fcd/skill-bake: hidden_act '{}' not supported yet (SiLU only)",
356 arch.hidden_act
357 ));
358 }
359 let has_linear = arch
360 .layer_types
361 .iter()
362 .any(|t| matches!(t, LayerType::LinearAttention));
363 let gdn = if has_linear {
364 let lc = arch
365 .linear_core
366 .as_ref()
367 .ok_or_else(|| "model has linear layers but no arch.linear_core".to_string())?;
368 if lc.kind != "gated_delta_net" {
369 return Err(format!(
370 "linear core '{}' has no FCD backward (only gated_delta_net)",
371 lc.kind
372 ));
373 }
374 Some(GdnDims {
375 nv: lc.num_heads,
376 nk: arch
377 .linear_num_key_heads
378 .ok_or("linear core needs arch.linear_num_key_heads")?,
379 dk: arch
380 .linear_key_head_dim
381 .ok_or("linear core needs arch.linear_key_head_dim")?,
382 dv: lc.value_head_dim,
383 kk: arch
384 .linear_conv_kernel_dim
385 .ok_or("linear core needs arch.linear_conv_kernel_dim")?,
386 })
387 } else {
388 None
389 };
390 let (nh, nkv, hd, h) = (
391 arch.num_attention_heads,
392 arch.num_kv_heads,
393 arch.head_dim,
394 arch.hidden_size,
395 );
396 let embed = deq(model, "model.embed_tokens.weight")?;
397 let lm_head = if model.tensor("lm_head.weight").is_some() {
398 Some(deq(model, "lm_head.weight")?)
399 } else if arch.tie_word_embeddings {
400 None
401 } else {
402 return Err("no lm_head.weight and tie_word_embeddings is false".into());
403 };
404 let final_norm = deq(model, "model.norm.weight")?;
405
406 let mut layers = Vec::with_capacity(arch.num_layers);
407 for li in 0..arch.num_layers {
408 let p = format!("model.layers.{li}.");
409 if model.tensor(&format!("{p}mlp.gate.weight")).is_some() {
410 return Err(format!(
411 "layer {li} is MoE — FCD polish supports dense FFN only"
412 ));
413 }
414 let attn = match arch.layer_types.get(li) {
415 Some(LayerType::LinearAttention) => {
416 let la = |n: &str| deq(model, &format!("{p}linear_attn.{n}"));
417 FcdAttn::Gdn {
418 wqkv: la("in_proj_qkv.weight")?,
419 wz: la("in_proj_z.weight")?,
420 wa: la("in_proj_a.weight")?,
421 wb: la("in_proj_b.weight")?,
422 conv: la("conv1d.weight")?,
423 a_log: la("A_log")?,
424 dt_bias: la("dt_bias")?,
425 norm: la("norm.weight")?,
426 wout: la("out_proj.weight")?,
427 }
428 }
429 _ => {
430 let wq = deq(model, &format!("{p}self_attn.q_proj.weight"))?;
431 let output_gate = wq.len() == 2 * nh * hd * h;
432 let opt = |n: &str| -> Option<Vec<f32>> {
433 model
434 .tensor(&format!("{p}self_attn.{n}"))
435 .and_then(|_| deq(model, &format!("{p}self_attn.{n}")).ok())
436 };
437 let bias = match (opt("q_proj.bias"), opt("k_proj.bias"), opt("v_proj.bias")) {
438 (Some(a), Some(b), Some(c)) => Some((a, b, c)),
439 _ => None,
440 };
441 let wk = deq(model, &format!("{p}self_attn.k_proj.weight"))?;
442 let wv = deq(model, &format!("{p}self_attn.v_proj.weight"))?;
443 let mut wqkv = Vec::with_capacity(wq.len() + wk.len() + wv.len());
444 wqkv.extend_from_slice(&wq);
445 wqkv.extend_from_slice(&wk);
446 wqkv.extend_from_slice(&wv);
447 FcdAttn::Full {
448 wq,
449 wk,
450 wv,
451 wqkv,
452 wo: deq(model, &format!("{p}self_attn.o_proj.weight"))?,
453 q_norm: opt("q_norm.weight"),
454 k_norm: opt("k_norm.weight"),
455 bias,
456 output_gate,
457 }
458 }
459 };
460 let gate = deq(model, &format!("{p}mlp.gate_proj.weight"))?;
461 let inter = gate.len() / h;
462 let up = deq(model, &format!("{p}mlp.up_proj.weight"))?;
463 let mut gu = Vec::with_capacity(gate.len() + up.len());
464 gu.extend_from_slice(&gate);
465 gu.extend_from_slice(&up);
466 layers.push(FcdLayer {
467 attn,
468 inter,
469 iln: deq(model, &format!("{p}input_layernorm.weight"))?,
470 pln: deq(model, &format!("{p}post_attention_layernorm.weight"))?,
471 gate,
472 up,
473 down: deq(model, &format!("{p}mlp.down_proj.weight"))?,
474 gu,
475 });
476 }
477
478 let rotary_dim = ((hd as f32 * arch.partial_rotary_factor) as usize)
479 .max(2)
480 .min(hd);
481 let base = arch.rope_theta;
482 let inv_freq: Vec<f64> = (0..rotary_dim / 2)
483 .map(|i| 1.0 / base.powf(2.0 * i as f64 / rotary_dim as f64))
484 .collect();
485 let loops = arch.num_loops.max(1);
486 let loop_norm = arch.loop_final_norm;
487 let mut flags = o1.layer_flags(arch.num_layers);
488 flags.resize(arch.num_layers, false);
489 for (li, f) in flags.iter_mut().enumerate() {
492 if *f && !matches!(layers[li].attn, FcdAttn::Full { .. }) {
493 *f = false;
494 }
495 }
496 Ok(Self {
497 hidden: h,
498 nh,
499 nkv,
500 hd,
501 nl: arch.num_layers,
502 vocab: arch.vocab_size.min(embed.len() / h),
503 eps: arch.rms_norm_eps,
504 gemma: matches!(arch.norm_style, NormStyle::Gemma),
505 rotary_dim,
506 inv_freq,
507 embed,
508 lm_head,
509 final_norm,
510 layers,
511 o1_flags: flags,
512 nys: NysCfg {
515 m: o1.m,
516 w: o1.w,
517 sink: o1.sink,
518 prefill: None,
519 },
520 gdn,
521 loops,
522 loop_norm,
523 pool: Pool::from_env(),
524 })
525 }
526
527 pub fn converted(&self) -> Vec<usize> {
529 (0..self.nl).filter(|&i| self.o1_flags[i]).collect()
530 }
531
532 fn head_weight(&self) -> &[f32] {
533 self.lm_head.as_deref().unwrap_or(&self.embed)
534 }
535}
536
537const PARAMS_PER_LAYER: usize = 5; pub struct TrainState {
544 pub layers: Vec<usize>,
545 pub data: Vec<Vec<f32>>,
547 grad: Vec<Vec<f32>>,
548 m1: Vec<Vec<f32>>,
549 m2: Vec<Vec<f32>>,
550 step_t: u64,
551}
552
553impl TrainState {
554 pub fn new(fm: &FcdModel) -> Self {
555 let layers = fm.converted();
556 let mut data = Vec::with_capacity(layers.len() * PARAMS_PER_LAYER);
557 for &li in &layers {
558 let l = &fm.layers[li];
559 data.push(l.iln.clone());
560 data.push(l.pln.clone());
561 data.push(l.gate.clone());
562 data.push(l.up.clone());
563 data.push(l.down.clone());
564 }
565 let zeros: Vec<Vec<f32>> = data.iter().map(|d| vec![0f32; d.len()]).collect();
566 Self {
567 layers,
568 grad: zeros.clone(),
569 m1: zeros.clone(),
570 m2: zeros,
571 data,
572 step_t: 0,
573 }
574 }
575
576 fn slot(&self, li: usize) -> Option<usize> {
577 self.layers.iter().position(|&x| x == li)
578 }
579
580 #[doc(hidden)]
582 pub fn grads(&self) -> &[Vec<f32>] {
583 &self.grad
584 }
585
586 fn zero_grad(&mut self) {
587 for g in &mut self.grad {
588 for v in g.iter_mut() {
589 *v = 0.0;
590 }
591 }
592 }
593
594 fn clip_and_step(&mut self, lr: f64) -> f64 {
597 let mut sq = 0f64;
598 for g in &self.grad {
599 for &v in g {
600 sq += (v as f64) * (v as f64);
601 }
602 }
603 let gn = sq.sqrt();
604 let scale = if gn > 1.0 { 1.0 / (gn + 1e-6) } else { 1.0 };
605 self.step_t += 1;
606 let bc1 = 1.0 - ADAM_B1.powi(self.step_t as i32);
607 let bc2 = 1.0 - ADAM_B2.powi(self.step_t as i32);
608 for p in 0..self.data.len() {
609 let (d, g, m, v) = (
610 &mut self.data[p],
611 &self.grad[p],
612 &mut self.m1[p],
613 &mut self.m2[p],
614 );
615 for i in 0..d.len() {
616 let gi = g[i] as f64 * scale;
617 let mi = ADAM_B1 * m[i] as f64 + (1.0 - ADAM_B1) * gi;
618 let vi = ADAM_B2 * v[i] as f64 + (1.0 - ADAM_B2) * gi * gi;
619 m[i] = mi as f32;
620 v[i] = vi as f32;
621 let upd = (mi / bc1) / ((vi / bc2).sqrt() + ADAM_EPS) + ADAM_WD * d[i] as f64;
622 d[i] = (d[i] as f64 - lr * upd) as f32;
623 }
624 }
625 gn
626 }
627}
628
629#[derive(Clone, Copy)]
632pub(crate) struct LnFfn<'a> {
633 pub(crate) iln: &'a [f32],
634 pub(crate) pln: &'a [f32],
635 pub(crate) gate: &'a [f32],
636 pub(crate) up: &'a [f32],
637 pub(crate) down: &'a [f32],
638 pub(crate) gu: Option<&'a [f32]>,
642}
643
644fn ln_ffn<'a>(fm: &'a FcdModel, ts: Option<&'a TrainState>, li: usize) -> LnFfn<'a> {
645 if let Some(t) = ts {
646 if let Some(s) = t.slot(li) {
647 let b = s * PARAMS_PER_LAYER;
648 return LnFfn {
649 iln: &t.data[b],
650 pln: &t.data[b + 1],
651 gate: &t.data[b + 2],
652 up: &t.data[b + 3],
653 down: &t.data[b + 4],
654 gu: None,
655 };
656 }
657 }
658 let l = &fm.layers[li];
659 LnFfn {
660 iln: &l.iln,
661 pln: &l.pln,
662 gate: &l.gate,
663 up: &l.up,
664 down: &l.down,
665 gu: Some(&l.gu),
666 }
667}
668
669enum AttnActs {
673 Full {
674 qpre: Vec<f32>,
675 kpre: Vec<f32>,
676 vproj: Vec<f32>,
677 qrot: Vec<f32>,
678 krot: Vec<f32>,
679 qinv: Vec<f32>,
680 kinv: Vec<f32>,
681 ao: Vec<f32>,
684 gate_pre: Vec<f32>,
686 },
687 Gdn {
690 qkv: Vec<f32>,
691 z: Vec<f32>,
692 a: Vec<f32>,
693 b: Vec<f32>,
694 },
695}
696
697pub(crate) struct LayerActs {
698 inv1: Vec<f32>,
699 attn: AttnActs,
700 pub(crate) h1: Vec<f32>,
701 pub(crate) n2: Vec<f32>,
702 pub(crate) inv2: Vec<f32>,
703 pub(crate) gpre: Vec<f32>,
704 pub(crate) upre: Vec<f32>,
705 pub(crate) act: Vec<f32>,
706}
707
708struct SendMut<T>(*mut T);
710unsafe impl<T> Send for SendMut<T> {}
711unsafe impl<T> Sync for SendMut<T> {}
712impl<T> SendMut<T> {
713 #[inline]
714 unsafe fn at(&self, i: usize) -> *mut T {
715 unsafe { self.0.add(i) }
716 }
717}
718
719impl FcdModel {
720 fn qk_norm_rope(
724 &self,
725 x: &mut [f32],
726 norm: Option<&[f32]>,
727 heads: usize,
728 t: usize,
729 inv_out: &mut [f32],
730 ) {
731 let hd = self.hd;
732 let n = x.len() / (heads * hd);
733 for r in 0..n {
734 let pos = r % t;
735 for hh in 0..heads {
736 let s = (r * heads + hh) * hd;
737 let head = &mut x[s..s + hd];
738 if let Some(w) = norm {
739 let mut inv = [0f32; 1];
740 let mut y = [0f32; 256];
741 debug_assert!(hd <= 256);
742 ops::rmsnorm_fwd(head, w, self.eps, self.gemma, &mut y[..hd], &mut inv);
743 head.copy_from_slice(&y[..hd]);
744 inv_out[r * heads + hh] = inv[0];
745 }
746 ops::rope_fwd(&mut head[..self.rotary_dim], pos, &self.inv_freq);
747 }
748 }
749 }
750
751 #[allow(clippy::too_many_arguments)]
757 fn layer_forward(
758 &self,
759 li: usize,
760 h_in: &[f32],
761 b: usize,
762 t: usize,
763 wts: &LnFfn,
764 nystrom: bool,
765 want_acts: bool,
766 ) -> (Vec<f32>, Option<LayerActs>) {
767 self.layer_forward_scaled(li, h_in, b, t, wts, nystrom, want_acts, None)
768 }
769
770 #[allow(clippy::too_many_arguments)]
774 pub(crate) fn layer_forward_scaled(
775 &self,
776 li: usize,
777 h_in: &[f32],
778 b: usize,
779 t: usize,
780 wts: &LnFfn,
781 nystrom: bool,
782 want_acts: bool,
783 ffn_scale: Option<&[f32]>,
784 ) -> (Vec<f32>, Option<LayerActs>) {
785 let hsz = self.hidden;
786 let n = b * t;
787 let l = &self.layers[li];
788 let pool = self.pool.as_deref();
789
790 let mut n1 = vec![0f32; n * hsz];
791 let mut inv1 = vec![0f32; n];
792 ops::rmsnorm_fwd(h_in, wts.iln, self.eps, self.gemma, &mut n1, &mut inv1);
793
794 let t_attn = std::time::Instant::now();
795 let (attn_out, attn_acts) = match &l.attn {
796 FcdAttn::Full { .. } => self.full_attn_fwd(&l.attn, &n1, b, t, nystrom),
797 FcdAttn::Gdn { .. } => self.gdn_attn_fwd(&l.attn, &n1, b, t),
798 };
799 prof::add(&prof::ATTN_FWD, t_attn);
800 let t_ffn = std::time::Instant::now();
801
802 let mut h1 = h_in.to_vec();
803 for (a, &x) in h1.iter_mut().zip(&attn_out) {
804 *a += x;
805 }
806
807 let mut n2 = vec![0f32; n * hsz];
808 let mut inv2 = vec![0f32; n];
809 ops::rmsnorm_fwd(&h1, wts.pln, self.eps, self.gemma, &mut n2, &mut inv2);
810
811 let inter = l.inter;
812 let mut gpre = vec![0f32; n * inter];
813 let mut upre = vec![0f32; n * inter];
814 if let Some(gu) = wts.gu {
815 let mut both = vec![0f32; n * 2 * inter];
818 ops::gemm_nt(&n2, gu, &mut both, n, hsz, 2 * inter, pool);
819 for r in 0..n {
820 let row = &both[r * 2 * inter..(r + 1) * 2 * inter];
821 gpre[r * inter..(r + 1) * inter].copy_from_slice(&row[..inter]);
822 upre[r * inter..(r + 1) * inter].copy_from_slice(&row[inter..]);
823 }
824 } else {
825 ops::gemm_nt(&n2, wts.gate, &mut gpre, n, hsz, inter, pool);
826 ops::gemm_nt(&n2, wts.up, &mut upre, n, hsz, inter, pool);
827 }
828 let mut act = vec![0f32; n * inter];
829 for i in 0..n * inter {
830 act[i] = ops::silu(gpre[i]) * upre[i];
831 }
832 let mut ffn = vec![0f32; n * hsz];
833 match ffn_scale {
834 Some(g) => {
835 debug_assert_eq!(g.len(), inter);
836 let mut act2 = act.clone();
837 for r in 0..n {
838 for (a, &gv) in act2[r * inter..(r + 1) * inter].iter_mut().zip(g) {
839 *a *= gv;
840 }
841 }
842 ops::gemm_nt(&act2, wts.down, &mut ffn, n, inter, hsz, pool);
843 }
844 None => ops::gemm_nt(&act, wts.down, &mut ffn, n, inter, hsz, pool),
845 }
846 let mut h2 = h1.clone();
847 for (a, &x) in h2.iter_mut().zip(&ffn) {
848 *a += x;
849 }
850
851 let acts = want_acts.then_some(LayerActs {
852 inv1,
853 attn: attn_acts,
854 h1,
855 n2,
856 inv2,
857 gpre,
858 upre,
859 act,
860 });
861 prof::add(&prof::FFN_FWD, t_ffn);
862 (h2, acts)
863 }
864
865 fn full_attn_fwd(
869 &self,
870 attn: &FcdAttn,
871 n1: &[f32],
872 b: usize,
873 t: usize,
874 nystrom: bool,
875 ) -> (Vec<f32>, AttnActs) {
876 let FcdAttn::Full {
877 wq,
878 wk,
879 wv,
880 wqkv,
881 wo,
882 q_norm,
883 k_norm,
884 bias,
885 output_gate,
886 } = attn
887 else {
888 unreachable!("full_attn_fwd on a non-Full layer");
889 };
890 let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
891 let n = b * t;
892 let pool = self.pool.as_deref();
893 let qdim = nh * hd;
894 let kvdim = nkv * hd;
895 let rep = nh / nkv;
896 let qrows = if *output_gate { 2 * qdim } else { qdim };
897
898 let fused = qrows + 2 * kvdim;
902 let mut qkv = vec![0f32; n * fused];
903 ops::gemm_nt(n1, wqkv, &mut qkv, n, hsz, fused, pool);
904 let _ = (wq, wk, wv);
905 let mut qraw = vec![0f32; n * qrows];
906 let mut kpre = vec![0f32; n * kvdim];
907 let mut vproj = vec![0f32; n * kvdim];
908 for r in 0..n {
909 let row = &qkv[r * fused..(r + 1) * fused];
910 qraw[r * qrows..(r + 1) * qrows].copy_from_slice(&row[..qrows]);
911 kpre[r * kvdim..(r + 1) * kvdim].copy_from_slice(&row[qrows..qrows + kvdim]);
912 vproj[r * kvdim..(r + 1) * kvdim].copy_from_slice(&row[qrows + kvdim..]);
913 }
914 if let Some((bq, bk, bv)) = bias {
915 for r in 0..n {
916 for (x, bb) in qraw[r * qrows..(r + 1) * qrows].iter_mut().zip(bq) {
917 *x += bb;
918 }
919 for (x, bb) in kpre[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bk) {
920 *x += bb;
921 }
922 for (x, bb) in vproj[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bv) {
923 *x += bb;
924 }
925 }
926 }
927 let (qpre, gate_pre) = if *output_gate {
929 let mut qh = vec![0f32; n * qdim];
930 let mut gp = vec![0f32; n * qdim];
931 for r in 0..n {
932 for h in 0..nh {
933 let src = r * qrows + 2 * h * hd;
934 let dst = r * qdim + h * hd;
935 qh[dst..dst + hd].copy_from_slice(&qraw[src..src + hd]);
936 gp[dst..dst + hd].copy_from_slice(&qraw[src + hd..src + 2 * hd]);
937 }
938 }
939 (qh, gp)
940 } else {
941 (qraw, Vec::new())
942 };
943
944 let mut qrot = qpre.clone();
945 let mut krot = kpre.clone();
946 let mut qinv = vec![0f32; n * nh];
947 let mut kinv = vec![0f32; n * nkv];
948 self.qk_norm_rope(&mut qrot, q_norm.as_deref(), nh, t, &mut qinv);
949 self.qk_norm_rope(&mut krot, k_norm.as_deref(), nkv, t, &mut kinv);
950
951 let mut ao = vec![0f32; n * qdim];
953 {
954 let units = b * nh;
955 let aop = SendMut(ao.as_mut_ptr());
956 let qr = &qrot;
957 let kr = &krot;
958 let vr = &vproj;
959 let nys = self.nys;
960 let run_unit = |u: usize| {
961 let (bi, h) = (u / nh, u % nh);
962 let g = h / rep;
963 if nystrom {
964 let mut q64 = vec![0f64; t * hd];
966 let mut k64 = vec![0f64; t * hd];
967 let mut v64 = vec![0f64; t * hd];
968 for p in 0..t {
969 let r = bi * t + p;
970 for c in 0..hd {
971 q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
972 k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
973 v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
974 }
975 }
976 let mut o64 = vec![0f64; t * hd];
977 ops::nystrom_head_fwd(&q64, &k64, &v64, t, hd, hd, &nys, &mut o64);
978 for p in 0..t {
979 let r = bi * t + p;
980 for c in 0..hd {
981 unsafe {
983 *aop.at(r * qdim + h * hd + c) = o64[p * hd + c] as f32;
984 }
985 }
986 }
987 } else {
988 let mut q32 = vec![0f32; t * hd];
989 let mut k32 = vec![0f32; t * hd];
990 let mut v32 = vec![0f32; t * hd];
991 for p in 0..t {
992 let r = bi * t + p;
993 q32[p * hd..(p + 1) * hd]
994 .copy_from_slice(&qr[r * qdim + h * hd..r * qdim + (h + 1) * hd]);
995 k32[p * hd..(p + 1) * hd]
996 .copy_from_slice(&kr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
997 v32[p * hd..(p + 1) * hd]
998 .copy_from_slice(&vr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
999 }
1000 let mut o32 = vec![0f32; t * hd];
1001 ops::attn_head_fwd(&q32, &k32, &v32, t, hd, hd, &mut o32);
1002 for p in 0..t {
1003 let r = bi * t + p;
1004 for c in 0..hd {
1005 unsafe {
1007 *aop.at(r * qdim + h * hd + c) = o32[p * hd + c];
1008 }
1009 }
1010 }
1011 }
1012 };
1013 match pool {
1014 Some(p) if units > 1 => p.run(&|widx, nw| {
1015 for u in (widx..units).step_by(nw) {
1016 run_unit(u);
1017 }
1018 }),
1019 _ => {
1020 for u in 0..units {
1021 run_unit(u);
1022 }
1023 }
1024 }
1025 }
1026
1027 let ao_eff: Vec<f32> = if *output_gate {
1029 ao.iter()
1030 .zip(&gate_pre)
1031 .map(|(&a, &g)| a * (1.0 / (1.0 + (-g).exp())))
1032 .collect()
1033 } else {
1034 ao.clone()
1035 };
1036 let mut attn_out = vec![0f32; n * hsz];
1037 ops::gemm_nt(&ao_eff, wo, &mut attn_out, n, qdim, hsz, pool);
1038 (
1039 attn_out,
1040 AttnActs::Full {
1041 qpre,
1042 kpre,
1043 vproj,
1044 qrot,
1045 krot,
1046 qinv,
1047 kinv,
1048 ao,
1049 gate_pre,
1050 },
1051 )
1052 }
1053
1054 fn gdn_attn_fwd(&self, attn: &FcdAttn, n1: &[f32], b: usize, t: usize) -> (Vec<f32>, AttnActs) {
1059 let FcdAttn::Gdn {
1060 wqkv,
1061 wz,
1062 wa,
1063 wb,
1064 conv,
1065 a_log,
1066 dt_bias,
1067 norm,
1068 wout,
1069 } = attn
1070 else {
1071 unreachable!("gdn_attn_fwd on a non-GDN layer");
1072 };
1073 let d = self.gdn.expect("gdn layer without gdn dims");
1074 let (hsz, n) = (self.hidden, b * t);
1075 let pool = self.pool.as_deref();
1076 let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
1077
1078 let mut qkv = vec![0f32; n * c_dim];
1079 ops::gemm_nt(n1, wqkv, &mut qkv, n, hsz, c_dim, pool);
1080 let mut z = vec![0f32; n * vd];
1081 ops::gemm_nt(n1, wz, &mut z, n, hsz, vd, pool);
1082 let mut a = vec![0f32; n * nv];
1083 ops::gemm_nt(n1, wa, &mut a, n, hsz, nv, pool);
1084 let mut bstr = vec![0f32; n * nv];
1085 ops::gemm_nt(n1, wb, &mut bstr, n, hsz, nv, pool);
1086
1087 let cfg = ops::GdnSeqCfg {
1088 nv: d.nv,
1089 nk: d.nk,
1090 dk: d.dk,
1091 dv: d.dv,
1092 kk: d.kk,
1093 rms_eps: self.eps,
1094 conv,
1095 a_log,
1096 dt_bias,
1097 norm,
1098 };
1099 let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
1101 let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
1102 let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
1103 let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
1104 let mut pre64 = vec![0f64; n * c_dim];
1105 let mut cq64 = vec![0f64; n * c_dim];
1106 for bi in 0..b {
1107 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1108 ops::gdn_conv_fwd(
1109 &qkv64[r.clone()],
1110 t,
1111 c_dim,
1112 d.kk,
1113 conv,
1114 &mut pre64[r.clone()],
1115 &mut cq64[r],
1116 );
1117 }
1118 let mut of = vec![0f32; n * vd];
1119 {
1120 let units = b * d.nk;
1121 let rep_v = d.nv / d.nk;
1122 let ofp = SendMut(of.as_mut_ptr());
1123 let (cqr, zr, ar, br) = (&cq64, &z64, &a64, &b64);
1124 let cfg_ref = &cfg;
1125 let run_unit = |u: usize| {
1126 let (bi, ko) = (u / d.nk, u % d.nk);
1127 let mut local = vec![0f64; t * vd];
1128 ops::gdn_group_fwd(
1129 &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
1130 &zr[bi * t * vd..(bi + 1) * t * vd],
1131 &ar[bi * t * nv..(bi + 1) * t * nv],
1132 &br[bi * t * nv..(bi + 1) * t * nv],
1133 t,
1134 cfg_ref,
1135 ko,
1136 &mut local,
1137 );
1138 for hh in 0..rep_v {
1139 let h = ko * rep_v + hh;
1140 for p in 0..t {
1141 for dj in 0..d.dv {
1142 unsafe {
1144 *ofp.at((bi * t + p) * vd + h * d.dv + dj) =
1145 local[p * vd + h * d.dv + dj] as f32;
1146 }
1147 }
1148 }
1149 }
1150 };
1151 match pool {
1152 Some(p) if units > 1 => p.run(&|widx, nw| {
1153 for u in (widx..units).step_by(nw) {
1154 run_unit(u);
1155 }
1156 }),
1157 _ => {
1158 for u in 0..units {
1159 run_unit(u);
1160 }
1161 }
1162 }
1163 }
1164 let mut attn_out = vec![0f32; n * hsz];
1165 ops::gemm_nt(&of, wout, &mut attn_out, n, vd, hsz, pool);
1166 (attn_out, AttnActs::Gdn { qkv, z, a, b: bstr })
1167 }
1168
1169 #[allow(clippy::too_many_arguments)]
1173 fn layer_backward(
1174 &self,
1175 li: usize,
1176 h_in: &[f32],
1177 b: usize,
1178 t: usize,
1179 wts: &LnFfn,
1180 nystrom: bool,
1181 acts: &LayerActs,
1182 dh2: &[f32],
1183 mut grads: Option<&mut [Vec<f32>]>,
1184 ) -> Vec<f32> {
1185 let hsz = self.hidden;
1186 let n = b * t;
1187 let l = &self.layers[li];
1188 let pool = self.pool.as_deref();
1189 let inter = l.inter;
1190
1191 let mut dact = vec![0f32; n * inter];
1193 ops::gemm_dx(dh2, wts.down, &mut dact, n, inter, hsz, pool);
1194 if let Some(g) = grads.as_deref_mut() {
1195 ops::gemm_dw(dh2, &acts.act, &mut g[4], n, inter, hsz, pool);
1196 }
1197 let mut dg = vec![0f32; n * inter];
1198 let mut du = vec![0f32; n * inter];
1199 for i in 0..n * inter {
1200 dg[i] = dact[i] * acts.upre[i] * ops::silu_bwd(acts.gpre[i]);
1201 du[i] = dact[i] * ops::silu(acts.gpre[i]);
1202 }
1203 let mut dn2 = vec![0f32; n * hsz];
1204 ops::gemm_dx(&dg, wts.gate, &mut dn2, n, hsz, inter, pool);
1205 ops::gemm_dx(&du, wts.up, &mut dn2, n, hsz, inter, pool);
1206 if let Some(g) = grads.as_deref_mut() {
1207 ops::gemm_dw(&dg, &acts.n2, &mut g[2], n, hsz, inter, pool);
1208 ops::gemm_dw(&du, &acts.n2, &mut g[3], n, hsz, inter, pool);
1209 }
1210
1211 let mut dh1 = dh2.to_vec();
1212 ops::rmsnorm_bwd(
1213 &acts.h1,
1214 wts.pln,
1215 &acts.inv2,
1216 &dn2,
1217 self.gemma,
1218 &mut dh1,
1219 grads.as_deref_mut().map(|g| &mut g[1][..]),
1220 );
1221
1222 let dn1 = match &l.attn {
1224 FcdAttn::Full { .. } => self.full_attn_bwd(&l.attn, &acts.attn, &dh1, b, t, nystrom),
1225 FcdAttn::Gdn { .. } => self.gdn_attn_bwd(&l.attn, &acts.attn, &dh1, b, t),
1226 };
1227
1228 let mut dh_in = dh1.clone();
1229 ops::rmsnorm_bwd(
1230 h_in,
1231 wts.iln,
1232 &acts.inv1,
1233 &dn1,
1234 self.gemma,
1235 &mut dh_in,
1236 grads.map(|g| &mut g[0][..]),
1237 );
1238 dh_in
1239 }
1240
1241 fn full_attn_bwd(
1245 &self,
1246 attn: &FcdAttn,
1247 acts: &AttnActs,
1248 dattn: &[f32],
1249 b: usize,
1250 t: usize,
1251 nystrom: bool,
1252 ) -> Vec<f32> {
1253 let FcdAttn::Full {
1254 wq,
1255 wqkv,
1256 wk,
1257 wv,
1258 wo,
1259 q_norm,
1260 k_norm,
1261 output_gate,
1262 ..
1263 } = attn
1264 else {
1265 unreachable!("full_attn_bwd on a non-Full layer");
1266 };
1267 let AttnActs::Full {
1268 qpre,
1269 kpre,
1270 vproj,
1271 qrot,
1272 krot,
1273 qinv,
1274 kinv,
1275 ao,
1276 gate_pre,
1277 } = acts
1278 else {
1279 unreachable!("acts mismatch");
1280 };
1281 let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
1282 let n = b * t;
1283 let pool = self.pool.as_deref();
1284 let qdim = nh * hd;
1285 let kvdim = nkv * hd;
1286 let rep = nh / nkv;
1287 let qrows = if *output_gate { 2 * qdim } else { qdim };
1288
1289 let mut dao_eff = vec![0f32; n * qdim];
1290 ops::gemm_dx(dattn, wo, &mut dao_eff, n, qdim, hsz, pool);
1291 let (dao, dgate) = if *output_gate {
1293 let mut dao = vec![0f32; n * qdim];
1294 let mut dgp = vec![0f32; n * qdim];
1295 for i in 0..n * qdim {
1296 let sig = 1.0 / (1.0 + (-gate_pre[i]).exp());
1297 dao[i] = dao_eff[i] * sig;
1298 dgp[i] = dao_eff[i] * ao[i] * sig * (1.0 - sig);
1299 }
1300 (dao, dgp)
1301 } else {
1302 (dao_eff, Vec::new())
1303 };
1304
1305 let mut dqrot = vec![0f32; n * qdim];
1306 let mut dkrot = vec![0f32; n * kvdim];
1307 let mut dvproj = vec![0f32; n * kvdim];
1308 {
1309 let units = b * nkv;
1312 let dqp = SendMut(dqrot.as_mut_ptr());
1313 let dkp = SendMut(dkrot.as_mut_ptr());
1314 let dvp = SendMut(dvproj.as_mut_ptr());
1315 let (qr, kr, vr) = (qrot, krot, vproj);
1316 let daor = &dao;
1317 let nys = self.nys;
1318 let run_unit = |u: usize| {
1319 let (bi, g) = (u / nkv, u % nkv);
1320 let mut k64 = vec![0f64; t * hd];
1321 let mut v64 = vec![0f64; t * hd];
1322 for p in 0..t {
1323 let r = bi * t + p;
1324 for c in 0..hd {
1325 k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
1326 v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
1327 }
1328 }
1329 let mut dk64 = vec![0f64; t * hd];
1330 let mut dv64 = vec![0f64; t * hd];
1331 let mut q64 = vec![0f64; t * hd];
1332 let mut do64 = vec![0f64; t * hd];
1333 let mut dq64 = vec![0f64; t * hd];
1334 for hh in 0..rep {
1335 let h = g * rep + hh;
1336 for p in 0..t {
1337 let r = bi * t + p;
1338 for c in 0..hd {
1339 q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
1340 do64[p * hd + c] = daor[r * qdim + h * hd + c] as f64;
1341 }
1342 }
1343 for v in dq64.iter_mut() {
1344 *v = 0.0;
1345 }
1346 if nystrom {
1347 ops::nystrom_head_bwd(
1348 &q64, &k64, &v64, &do64, t, hd, hd, &nys, &mut dq64, &mut dk64,
1349 &mut dv64,
1350 );
1351 } else {
1352 ops::attn_head_bwd(
1353 &q64, &k64, &v64, &do64, t, hd, hd, &mut dq64, &mut dk64, &mut dv64,
1354 );
1355 }
1356 for p in 0..t {
1357 let r = bi * t + p;
1358 for c in 0..hd {
1359 unsafe {
1361 *dqp.at(r * qdim + h * hd + c) = dq64[p * hd + c] as f32;
1362 }
1363 }
1364 }
1365 }
1366 for p in 0..t {
1367 let r = bi * t + p;
1368 for c in 0..hd {
1369 unsafe {
1371 *dkp.at(r * kvdim + g * hd + c) = dk64[p * hd + c] as f32;
1372 *dvp.at(r * kvdim + g * hd + c) = dv64[p * hd + c] as f32;
1373 }
1374 }
1375 }
1376 };
1377 match pool {
1378 Some(p) if units > 1 => p.run(&|widx, nw| {
1379 for u in (widx..units).step_by(nw) {
1380 run_unit(u);
1381 }
1382 }),
1383 _ => {
1384 for u in 0..units {
1385 run_unit(u);
1386 }
1387 }
1388 }
1389 }
1390
1391 let mut dqpre = vec![0f32; n * qdim];
1393 let mut dkpre = vec![0f32; n * kvdim];
1394 for r in 0..n {
1395 let pos = r % t;
1396 for h in 0..nh {
1397 let s = r * qdim + h * hd;
1398 ops::rope_bwd(&mut dqrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1399 match q_norm {
1400 Some(w) => ops::rmsnorm_bwd(
1401 &qpre[s..s + hd],
1402 w,
1403 &qinv[r * nh + h..r * nh + h + 1],
1404 &dqrot[s..s + hd],
1405 self.gemma,
1406 &mut dqpre[s..s + hd],
1407 None,
1408 ),
1409 None => dqpre[s..s + hd].copy_from_slice(&dqrot[s..s + hd]),
1410 }
1411 }
1412 for g in 0..nkv {
1413 let s = r * kvdim + g * hd;
1414 ops::rope_bwd(&mut dkrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1415 match k_norm {
1416 Some(w) => ops::rmsnorm_bwd(
1417 &kpre[s..s + hd],
1418 w,
1419 &kinv[r * nkv + g..r * nkv + g + 1],
1420 &dkrot[s..s + hd],
1421 self.gemma,
1422 &mut dkpre[s..s + hd],
1423 None,
1424 ),
1425 None => dkpre[s..s + hd].copy_from_slice(&dkrot[s..s + hd]),
1426 }
1427 }
1428 }
1429
1430 let dqraw: Vec<f32> = if *output_gate {
1432 let mut dq = vec![0f32; n * qrows];
1433 for r in 0..n {
1434 for h in 0..nh {
1435 let dst = r * qrows + 2 * h * hd;
1436 let src = r * qdim + h * hd;
1437 dq[dst..dst + hd].copy_from_slice(&dqpre[src..src + hd]);
1438 dq[dst + hd..dst + 2 * hd].copy_from_slice(&dgate[src..src + hd]);
1439 }
1440 }
1441 dq
1442 } else {
1443 dqpre
1444 };
1445
1446 let fused = qrows + 2 * kvdim;
1451 let mut dqkv = vec![0f32; n * fused];
1452 for r in 0..n {
1453 let row = &mut dqkv[r * fused..(r + 1) * fused];
1454 row[..qrows].copy_from_slice(&dqraw[r * qrows..(r + 1) * qrows]);
1455 row[qrows..qrows + kvdim].copy_from_slice(&dkpre[r * kvdim..(r + 1) * kvdim]);
1456 row[qrows + kvdim..].copy_from_slice(&dvproj[r * kvdim..(r + 1) * kvdim]);
1457 }
1458 let mut dn1 = vec![0f32; n * hsz];
1459 ops::gemm_dx(&dqkv, wqkv, &mut dn1, n, hsz, fused, pool);
1460 let _ = (wq, wk, wv);
1461 dn1
1462 }
1463
1464 fn gdn_attn_bwd(
1468 &self,
1469 attn: &FcdAttn,
1470 acts: &AttnActs,
1471 dattn: &[f32],
1472 b: usize,
1473 t: usize,
1474 ) -> Vec<f32> {
1475 let FcdAttn::Gdn {
1476 wqkv,
1477 wz,
1478 wa,
1479 wb,
1480 conv,
1481 a_log,
1482 dt_bias,
1483 norm,
1484 wout,
1485 } = attn
1486 else {
1487 unreachable!("gdn_attn_bwd on a non-GDN layer");
1488 };
1489 let AttnActs::Gdn { qkv, z, a, b: bstr } = acts else {
1490 unreachable!("acts mismatch");
1491 };
1492 let d = self.gdn.expect("gdn layer without gdn dims");
1493 let (hsz, n) = (self.hidden, b * t);
1494 let pool = self.pool.as_deref();
1495 let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
1496
1497 let mut dof = vec![0f32; n * vd];
1498 ops::gemm_dx(dattn, wout, &mut dof, n, vd, hsz, pool);
1499
1500 let cfg = ops::GdnSeqCfg {
1501 nv: d.nv,
1502 nk: d.nk,
1503 dk: d.dk,
1504 dv: d.dv,
1505 kk: d.kk,
1506 rms_eps: self.eps,
1507 conv,
1508 a_log,
1509 dt_bias,
1510 norm,
1511 };
1512 let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
1513 let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
1514 let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
1515 let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
1516 let dof64: Vec<f64> = dof.iter().map(|&v| v as f64).collect();
1517 let mut pre64 = vec![0f64; n * c_dim];
1518 let mut cq64 = vec![0f64; n * c_dim];
1519 for bi in 0..b {
1520 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1521 ops::gdn_conv_fwd(
1522 &qkv64[r.clone()],
1523 t,
1524 c_dim,
1525 d.kk,
1526 conv,
1527 &mut pre64[r.clone()],
1528 &mut cq64[r],
1529 );
1530 }
1531
1532 let mut dcq64 = vec![0f64; n * c_dim];
1533 let mut dz64 = vec![0f64; n * vd];
1534 let mut da64 = vec![0f64; n * nv];
1535 let mut db64 = vec![0f64; n * nv];
1536 {
1537 let units = b * d.nk;
1538 let rep_v = d.nv / d.nk;
1539 let kd = d.nk * d.dk;
1540 let dcqp = SendMut(dcq64.as_mut_ptr());
1541 let dzp = SendMut(dz64.as_mut_ptr());
1542 let dap = SendMut(da64.as_mut_ptr());
1543 let dbp = SendMut(db64.as_mut_ptr());
1544 let (cqr, zr, ar, br, dor) = (&cq64, &z64, &a64, &b64, &dof64);
1545 let cfg_ref = &cfg;
1546 let run_unit = |u: usize| {
1547 let (bi, ko) = (u / d.nk, u % d.nk);
1548 let mut dcq_l = vec![0f64; t * c_dim];
1551 let mut dz_l = vec![0f64; t * vd];
1552 let mut da_l = vec![0f64; t * nv];
1553 let mut db_l = vec![0f64; t * nv];
1554 ops::gdn_group_bwd(
1555 &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
1556 &zr[bi * t * vd..(bi + 1) * t * vd],
1557 &ar[bi * t * nv..(bi + 1) * t * nv],
1558 &br[bi * t * nv..(bi + 1) * t * nv],
1559 t,
1560 cfg_ref,
1561 ko,
1562 &dor[bi * t * vd..(bi + 1) * t * vd],
1563 &mut dcq_l,
1564 &mut dz_l,
1565 &mut da_l,
1566 &mut db_l,
1567 );
1568 for p in 0..t {
1571 let row = (bi * t + p) * c_dim;
1572 for c in ko * d.dk..(ko + 1) * d.dk {
1573 unsafe {
1574 *dcqp.at(row + c) = dcq_l[p * c_dim + c];
1575 *dcqp.at(row + kd + c) = dcq_l[p * c_dim + kd + c];
1576 }
1577 }
1578 for hh in 0..rep_v {
1579 let h = ko * rep_v + hh;
1580 for dj in 0..d.dv {
1581 unsafe {
1582 *dcqp.at(row + 2 * kd + h * d.dv + dj) =
1583 dcq_l[p * c_dim + 2 * kd + h * d.dv + dj];
1584 *dzp.at((bi * t + p) * vd + h * d.dv + dj) =
1585 dz_l[p * vd + h * d.dv + dj];
1586 }
1587 }
1588 unsafe {
1589 *dap.at((bi * t + p) * nv + h) = da_l[p * nv + h];
1590 *dbp.at((bi * t + p) * nv + h) = db_l[p * nv + h];
1591 }
1592 }
1593 }
1594 };
1595 match pool {
1596 Some(p) if units > 1 => p.run(&|widx, nw| {
1597 for u in (widx..units).step_by(nw) {
1598 run_unit(u);
1599 }
1600 }),
1601 _ => {
1602 for u in 0..units {
1603 run_unit(u);
1604 }
1605 }
1606 }
1607 }
1608
1609 let mut dqkv64 = vec![0f64; n * c_dim];
1610 for bi in 0..b {
1611 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1612 ops::gdn_conv_bwd(
1613 &pre64[r.clone()],
1614 t,
1615 c_dim,
1616 d.kk,
1617 conv,
1618 &dcq64[r.clone()],
1619 &mut dqkv64[r],
1620 );
1621 }
1622 let to32 = |v: &[f64]| -> Vec<f32> { v.iter().map(|&x| x as f32).collect() };
1623 let (dqkv, dz, da, db) = (to32(&dqkv64), to32(&dz64), to32(&da64), to32(&db64));
1624
1625 let mut dn1 = vec![0f32; n * hsz];
1626 ops::gemm_dx(&dqkv, wqkv, &mut dn1, n, hsz, c_dim, pool);
1627 ops::gemm_dx(&dz, wz, &mut dn1, n, hsz, vd, pool);
1628 ops::gemm_dx(&da, wa, &mut dn1, n, hsz, nv, pool);
1629 ops::gemm_dx(&db, wb, &mut dn1, n, hsz, nv, pool);
1630 dn1
1631 }
1632
1633 fn forward_hidden(
1638 &self,
1639 ids: &[u32],
1640 b: usize,
1641 t: usize,
1642 ts: Option<&TrainState>,
1643 student: bool,
1644 mut keep: Option<&mut Vec<Vec<f32>>>,
1645 ) -> Vec<f32> {
1646 let hsz = self.hidden;
1647 let mut h = vec![0f32; b * t * hsz];
1648 for (r, &id) in ids.iter().enumerate() {
1649 let src = (id as usize).min(self.embed.len() / hsz - 1) * hsz;
1650 h[r * hsz..(r + 1) * hsz].copy_from_slice(&self.embed[src..src + hsz]);
1651 }
1652 for li in 0..self.nl {
1653 if let Some(k) = keep.as_deref_mut() {
1654 k.push(h.clone());
1655 }
1656 let wts = ln_ffn(self, if student { ts } else { None }, li);
1657 let nys = student && self.o1_flags[li];
1658 h = self.layer_forward(li, &h, b, t, &wts, nys, false).0;
1659 }
1660 h
1661 }
1662
1663 fn loss_and_dhidden(
1666 &self,
1667 hs: &[f32],
1668 ht: &[f32],
1669 targets: &[u32],
1670 kl_w: f64,
1671 ) -> (f64, f64, Vec<f32>) {
1672 let hsz = self.hidden;
1673 let n = targets.len();
1674 let pool = self.pool.as_deref();
1675 let wh = self.head_weight();
1676 let vs = self.vocab;
1677
1678 let mut ns = vec![0f32; n * hsz];
1679 let mut invs = vec![0f32; n];
1680 ops::rmsnorm_fwd(
1681 hs,
1682 &self.final_norm,
1683 self.eps,
1684 self.gemma,
1685 &mut ns,
1686 &mut invs,
1687 );
1688 let mut nt = vec![0f32; n * hsz];
1689 let mut invt = vec![0f32; n];
1690 ops::rmsnorm_fwd(
1691 ht,
1692 &self.final_norm,
1693 self.eps,
1694 self.gemma,
1695 &mut nt,
1696 &mut invt,
1697 );
1698
1699 let inv_n = 1.0 / n as f64;
1700 let mut ce_sum = 0f64;
1701 let mut kl_sum = 0f64;
1702 let mut dns = vec![0f32; n * hsz];
1703 let mut ls = vec![0f32; LM_CHUNK * vs];
1704 let mut lt = vec![0f32; LM_CHUNK * vs];
1705 let mut dlg = vec![0f32; LM_CHUNK * vs];
1706 let mut r0 = 0usize;
1707 while r0 < n {
1708 let r1 = (r0 + LM_CHUNK).min(n);
1709 let c = r1 - r0;
1710 ops::gemm_nt(
1711 &ns[r0 * hsz..r1 * hsz],
1712 wh,
1713 &mut ls[..c * vs],
1714 c,
1715 hsz,
1716 vs,
1717 pool,
1718 );
1719 ops::gemm_nt(
1720 &nt[r0 * hsz..r1 * hsz],
1721 wh,
1722 &mut lt[..c * vs],
1723 c,
1724 hsz,
1725 vs,
1726 pool,
1727 );
1728 for r in 0..c {
1729 let (ce, kl) = ops::ce_kl_position(
1730 &ls[r * vs..(r + 1) * vs],
1731 <[r * vs..(r + 1) * vs],
1732 targets[r0 + r] as usize,
1733 kl_w,
1734 inv_n,
1735 &mut dlg[r * vs..(r + 1) * vs],
1736 );
1737 ce_sum += ce;
1738 kl_sum += kl;
1739 }
1740 ops::gemm_dx(
1741 &dlg[..c * vs],
1742 wh,
1743 &mut dns[r0 * hsz..r1 * hsz],
1744 c,
1745 hsz,
1746 vs,
1747 pool,
1748 );
1749 r0 = r1;
1750 }
1751
1752 let mut dhs = vec![0f32; n * hsz];
1753 ops::rmsnorm_bwd(
1754 hs,
1755 &self.final_norm,
1756 &invs,
1757 &dns,
1758 self.gemma,
1759 &mut dhs,
1760 None,
1761 );
1762 (ce_sum * inv_n, kl_sum * inv_n, dhs)
1763 }
1764
1765 fn backward(
1768 &self,
1769 b: usize,
1770 t: usize,
1771 keep: &[Vec<f32>],
1772 dh_last: Vec<f32>,
1773 ts: &mut TrainState,
1774 ) {
1775 let TrainState {
1778 layers, data, grad, ..
1779 } = ts;
1780 let mut dh = dh_last;
1781 for li in (0..self.nl).rev() {
1782 let h_in = &keep[li];
1783 let nys = self.o1_flags[li];
1784 let slot = layers.iter().position(|&x| x == li);
1785 let wts = match slot {
1786 Some(s) => {
1787 let bi = s * PARAMS_PER_LAYER;
1788 LnFfn {
1789 iln: &data[bi],
1790 pln: &data[bi + 1],
1791 gate: &data[bi + 2],
1792 up: &data[bi + 3],
1793 down: &data[bi + 4],
1794 gu: None,
1795 }
1796 }
1797 None => {
1798 let l = &self.layers[li];
1799 LnFfn {
1800 iln: &l.iln,
1801 pln: &l.pln,
1802 gate: &l.gate,
1803 up: &l.up,
1804 down: &l.down,
1805 gu: Some(&l.gu),
1806 }
1807 }
1808 };
1809 let (_, acts) = self.layer_forward(li, h_in, b, t, &wts, nys, true);
1810 let acts = acts.expect("want_acts");
1811 dh = match slot {
1812 Some(s) => {
1813 let gb = s * PARAMS_PER_LAYER;
1814 let gr = &mut grad[gb..gb + PARAMS_PER_LAYER];
1815 self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, Some(gr))
1816 }
1817 None => self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, None),
1818 };
1819 }
1820 }
1821
1822 #[doc(hidden)]
1829 pub fn loss_and_grads_for_test(
1830 &self,
1831 ids: &[u32],
1832 tgt: &[u32],
1833 b: usize,
1834 t: usize,
1835 ts: &mut TrainState,
1836 kl_w: f64,
1837 ) -> f64 {
1838 let ht = self.forward_hidden(ids, b, t, None, false, None);
1839 let mut keep = Vec::with_capacity(self.nl);
1840 let hs = self.forward_hidden(ids, b, t, Some(ts), true, Some(&mut keep));
1841 let (ce, kl, dhs) = self.loss_and_dhidden(&hs, &ht, tgt, kl_w);
1842 ts.zero_grad();
1843 self.backward(b, t, &keep, dhs, ts);
1844 (1.0 - kl_w) * ce + kl_w * kl
1845 }
1846
1847 pub fn val_ppl(
1851 &self,
1852 va: &[u32],
1853 ts: Option<&TrainState>,
1854 student: bool,
1855 bs: usize,
1856 nrounds: usize,
1857 seq: usize,
1858 ) -> f64 {
1859 let nwin = nrounds * bs;
1860 if va.len() < seq + 2 || nwin == 0 {
1861 return f64::NAN;
1862 }
1863 let stride = (va.len() - seq - 1) / nwin;
1864 let hsz = self.hidden;
1865 let wh = self.head_weight();
1866 let vs = self.vocab;
1867 let pool = self.pool.as_deref();
1868 let mut nll = 0f64;
1869 let mut cnt = 0usize;
1870 for j in 0..nrounds {
1871 let mut ids = Vec::with_capacity(bs * seq);
1872 let mut tgt = Vec::with_capacity(bs * seq);
1873 for bi in 0..bs {
1874 let off = ((j * bs + bi) * stride.max(1)).min(va.len() - seq - 1);
1875 ids.extend_from_slice(&va[off..off + seq]);
1876 tgt.extend_from_slice(&va[off + 1..off + seq + 1]);
1877 }
1878 let h = self.forward_hidden(&ids, bs, seq, ts, student, None);
1879 let n = bs * seq;
1880 let mut ns = vec![0f32; n * hsz];
1881 let mut inv = vec![0f32; n];
1882 ops::rmsnorm_fwd(
1883 &h,
1884 &self.final_norm,
1885 self.eps,
1886 self.gemma,
1887 &mut ns,
1888 &mut inv,
1889 );
1890 let mut lg = vec![0f32; LM_CHUNK * vs];
1891 let mut r0 = 0usize;
1892 while r0 < n {
1893 let r1 = (r0 + LM_CHUNK).min(n);
1894 let c = r1 - r0;
1895 ops::gemm_nt(
1896 &ns[r0 * hsz..r1 * hsz],
1897 wh,
1898 &mut lg[..c * vs],
1899 c,
1900 hsz,
1901 vs,
1902 pool,
1903 );
1904 for r in 0..c {
1905 let row = &lg[r * vs..(r + 1) * vs];
1906 let target = tgt[r0 + r] as usize;
1907 let mut mx = f64::NEG_INFINITY;
1908 for &v in row {
1909 mx = mx.max(v as f64);
1910 }
1911 let mut s = 0f64;
1912 for &v in row {
1913 s += (v as f64 - mx).exp();
1914 }
1915 nll += mx + s.ln() - row[target.min(vs - 1)] as f64;
1916 cnt += 1;
1917 }
1918 r0 = r1;
1919 }
1920 }
1921 (nll / cnt.max(1) as f64).exp()
1922 }
1923}
1924
1925pub fn run_polish(
1937 model: &Arc<CmfModel>,
1938 o1: &O1Cfg,
1939 hp: &FcdHyper,
1940 tr: &[u32],
1941 va: &[u32],
1942 out: &std::path::Path,
1943 gate: Option<&GenGateCfg>,
1944) -> Result<FcdReport, String> {
1945 if tr.len() < hp.seq + 2 {
1946 return Err(format!(
1947 "train corpus too small: {} tokens < seq+2 = {}",
1948 tr.len(),
1949 hp.seq + 2
1950 ));
1951 }
1952 let fm = FcdModel::from_cmf(model, o1)?;
1953 let converted = fm.converted();
1954 if converted.is_empty() {
1955 return Err("no converted layers under this --o1 spec (nothing to polish)".into());
1956 }
1957 tracing::info!(
1958 "fcd: {} layers converted ({} trainable tensors), m={} w={} sink={}, \
1959 corpus train {} / val {} tokens",
1960 converted.len(),
1961 converted.len() * PARAMS_PER_LAYER,
1962 fm.nys.m,
1963 fm.nys.w,
1964 fm.nys.sink,
1965 tr.len(),
1966 va.len()
1967 );
1968
1969 let mut ts = TrainState::new(&fm);
1970 let teacher_ppl = fm.val_ppl(va, None, false, hp.bs, 2, hp.seq);
1971 let ppl_start = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
1972 tracing::info!(
1973 "fcd: quick-val teacher ppl {teacher_ppl:.2} | zero-shot o1 student ppl {ppl_start:.2}"
1974 );
1975
1976 let mut gate_state: Option<(Pipeline, Vec<f64>)> = match gate {
1978 Some(g) if !g.prompts.is_empty() => {
1979 let greedy = SamplerConfig {
1980 temperature: 0.0,
1981 top_p: 1.0,
1982 top_k: 0,
1983 repetition_penalty: 1.0,
1984 min_p: 0.0,
1985 seed: Some(0),
1986 suppress_tokens: Vec::new(),
1987 };
1988 let mut pipe = Pipeline::from_model(model, greedy)
1989 .map_err(|e| format!("gen-gate pipeline: {e}"))?;
1990 pipe.set_o1(Some(o1.clone()));
1991 apply_trainables(&mut pipe, &fm, &ts);
1992 let base = gate_gen_scores(&mut pipe, g)?;
1993 tracing::info!("fcd gen-gate baseline loop-scores: {base:?}");
1994 Some((pipe, base))
1995 }
1996 Some(_) => {
1997 tracing::warn!("fcd gen-gate requested but val stream too short — gate off");
1998 None
1999 }
2000 None => None,
2001 };
2002 let init_snapshot: Option<Vec<Vec<f32>>> = gate_state.is_some().then(|| ts.data.clone());
2004 let mut gate_evals: Vec<(usize, f64, Vec<f64>, bool)> = Vec::new();
2005
2006 let mut rng = SplitMix64::new(hp.seed);
2007 let mut best: (f64, Option<Vec<Vec<f32>>>, usize) = (f64::INFINITY, None, 0);
2008 let mut losses: Vec<(f64, f64)> = Vec::with_capacity(hp.steps);
2009 let t0 = std::time::Instant::now();
2010 let n_per_step = hp.bs * hp.seq;
2011 for st in 1..=hp.steps {
2012 let mut ids = Vec::with_capacity(n_per_step);
2015 let mut tgt = Vec::with_capacity(n_per_step);
2016 for _ in 0..hp.bs {
2017 let off = (rng.next_u64() as usize) % (tr.len() - hp.seq - 1);
2018 ids.extend_from_slice(&tr[off..off + hp.seq]);
2019 tgt.extend_from_slice(&tr[off + 1..off + hp.seq + 1]);
2020 }
2021
2022 let ht = fm.forward_hidden(&ids, hp.bs, hp.seq, None, false, None);
2023 let mut keep: Vec<Vec<f32>> = Vec::with_capacity(fm.nl);
2024 let hs = fm.forward_hidden(&ids, hp.bs, hp.seq, Some(&ts), true, Some(&mut keep));
2025 let (ce, kl, dhs) = fm.loss_and_dhidden(&hs, &ht, &tgt, hp.kl_w);
2026 ts.zero_grad();
2027 fm.backward(hp.bs, hp.seq, &keep, dhs, &mut ts);
2028 let gn = ts.clip_and_step(hp.lr);
2029 losses.push((ce, kl));
2030
2031 let el = t0.elapsed().as_secs_f64();
2032 tracing::info!(
2033 "fcd step {st}/{}: ce {ce:.3} kl {kl:.3} |g| {gn:.3} ({:.1}s/step)",
2034 hp.steps,
2035 el / st as f64
2036 );
2037 if hp.eval_every > 0 && st % hp.eval_every == 0 {
2038 let p = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
2039 match (&mut gate_state, gate) {
2040 (Some((pipe, base)), Some(g)) => {
2041 apply_trainables(pipe, &fm, &ts);
2042 let scores = gate_gen_scores(pipe, g)?;
2043 let pass = gate_pass(&scores, base, g.threshold, g.baseline_slack);
2044 let tag = if pass && p < best.0 {
2045 best = (p, Some(ts.data.clone()), st);
2046 " *best*"
2047 } else {
2048 ""
2049 };
2050 tracing::info!(
2051 "fcd eval step {st}: val ppl {p:.2} | gen-gate {} (loop-scores {scores:?}){tag}",
2052 if pass { "PASS" } else { "FAIL" }
2053 );
2054 gate_evals.push((st, p, scores, pass));
2055 }
2056 _ => {
2057 let tag = if p < best.0 {
2058 best = (p, Some(ts.data.clone()), st);
2059 " *best*"
2060 } else {
2061 ""
2062 };
2063 tracing::info!("fcd eval step {st}: val ppl {p:.2}{tag}");
2064 }
2065 }
2066 }
2067 }
2068
2069 let mut gate_chosen: Option<usize> = None;
2073 if let Some(snap) = best.1.take() {
2074 ts.data = snap;
2075 gate_chosen = Some(best.2);
2076 tracing::info!(
2077 "fcd: restored best checkpoint from step {} (val ppl {:.2})",
2078 best.2,
2079 best.0
2080 );
2081 } else if let Some(init) = init_snapshot {
2082 ts.data = init;
2083 tracing::info!(
2084 "fcd: polish rejected by generation gate — identity artifact (zero-shot state written; claim 13 floor)"
2085 );
2086 }
2087 let ppl_final = fm.val_ppl(va, Some(&ts), true, hp.bs, 6, hp.seq);
2088 let report = FcdReport {
2089 converted: converted.clone(),
2090 teacher_ppl,
2091 ppl_start,
2092 ppl_best: best.0.min(ppl_final),
2093 best_step: best.2,
2094 ppl_final,
2095 steps_run: hp.steps,
2096 sec_per_step: t0.elapsed().as_secs_f64() / hp.steps.max(1) as f64,
2097 losses,
2098 gate: gate_state.map(|(_, base)| GateReport {
2099 baseline: base,
2100 evals: gate_evals,
2101 chosen: gate_chosen,
2102 }),
2103 };
2104 save_polished(model, out, &fm, &ts, o1, hp, &report)?;
2105 Ok(report)
2106}
2107
2108fn apply_trainables(pipe: &mut Pipeline, fm: &FcdModel, ts: &TrainState) {
2112 let hidden = fm.hidden;
2113 for (slot, &li) in ts.layers.iter().enumerate() {
2114 let b = slot * PARAMS_PER_LAYER;
2115 let inter = fm.layers[li].inter;
2116 let lw = &mut pipe.weights.layers[li];
2117 lw.input_norm = ts.data[b].clone();
2118 lw.post_norm = ts.data[b + 1].clone();
2119 lw.ffn = FfnKind::Dense(DenseFfn {
2120 gate_proj: QTensor::from_f32(ts.data[b + 2].clone(), inter, hidden),
2121 up_proj: QTensor::from_f32(ts.data[b + 3].clone(), inter, hidden),
2122 down_proj: QTensor::from_f32(ts.data[b + 4].clone(), hidden, inter),
2123 act: crate::pipeline::Act::Silu,
2124 });
2125 }
2126}
2127
2128fn gate_gen_scores(pipe: &mut Pipeline, g: &GenGateCfg) -> Result<Vec<f64>, String> {
2130 g.prompts
2131 .iter()
2132 .map(|p| {
2133 pipe.generate_from_ids(p, g.gen_tokens, None, None)
2134 .map(|r| loop_score(&r.token_ids))
2135 })
2136 .collect()
2137}
2138
2139fn save_polished(
2144 model: &CmfModel,
2145 out: &std::path::Path,
2146 fm: &FcdModel,
2147 ts: &TrainState,
2148 o1: &O1Cfg,
2149 hp: &FcdHyper,
2150 report: &FcdReport,
2151) -> Result<(), String> {
2152 use cortiq_core::format::TensorSpec;
2153 let mut replace: std::collections::HashMap<String, (usize, usize)> =
2154 std::collections::HashMap::new(); for (s, &li) in ts.layers.iter().enumerate() {
2156 let p = format!("model.layers.{li}.");
2157 for (k, suffix) in [
2158 (0usize, "input_layernorm.weight"),
2159 (1, "post_attention_layernorm.weight"),
2160 (2, "mlp.gate_proj.weight"),
2161 (3, "mlp.up_proj.weight"),
2162 (4, "mlp.down_proj.weight"),
2163 ] {
2164 replace.insert(format!("{p}{suffix}"), (s, k));
2165 }
2166 }
2167 let mut specs = Vec::with_capacity(model.tensors.len());
2168 for t in &model.tensors {
2169 if let Some(&(s, k)) = replace.get(&t.name) {
2170 let data = &ts.data[s * PARAMS_PER_LAYER + k];
2171 let mut bytes = Vec::with_capacity(data.len() * 4);
2172 for v in data {
2173 bytes.extend_from_slice(&v.to_le_bytes());
2174 }
2175 specs.push(TensorSpec {
2176 name: t.name.clone(),
2177 dtype: TensorDtype::F32,
2178 shape: t.shape.clone(),
2179 data: bytes,
2180 });
2181 } else {
2182 specs.push(TensorSpec {
2183 name: t.name.clone(),
2184 dtype: t.dtype,
2185 shape: t.shape.clone(),
2186 data: model.entry_bytes(t).to_vec(),
2187 });
2188 }
2189 }
2190
2191 let mut header = model.header.clone();
2192 let mut prov = match header.provenance.take() {
2193 Some(serde_json::Value::Object(m)) => m,
2194 _ => serde_json::Map::new(),
2195 };
2196 let layers_json = match &o1.layers {
2197 O1Layers::All => serde_json::json!("all"),
2198 O1Layers::Deep(n) => serde_json::json!(format!("deep{n}")),
2199 O1Layers::List(v) => serde_json::json!(v),
2200 };
2201 prov.insert(
2202 "o1_attn".into(),
2203 serde_json::json!({
2204 "layers": layers_json, "m": o1.m, "w": o1.w, "sink": o1.sink
2205 }),
2206 );
2207 prov.insert(
2208 "fcd".into(),
2209 serde_json::json!({
2210 "steps": hp.steps, "lr": hp.lr, "kl_w": hp.kl_w,
2211 "bs": hp.bs, "seq": hp.seq,
2212 "teacher_ppl": report.teacher_ppl,
2213 "ppl_start": report.ppl_start,
2214 "ppl_final": report.ppl_final,
2215 "best_step": report.best_step,
2216 "converted_layers": report.converted,
2217 }),
2218 );
2219 header.provenance = Some(serde_json::Value::Object(prov));
2220 let _ = fm; let masks = if model.masks.masks.is_empty() {
2223 None
2224 } else {
2225 Some(&model.masks)
2226 };
2227 CmfModel::write(out, &header, &specs, masks, model.vocab.as_deref())
2228 .map_err(|e| format!("writing polished cmf: {e}"))
2229}
2230
2231#[cfg(test)]
2232mod tests {
2233 use super::*;
2234
2235 #[test]
2237 fn gate_selects_lowest_ppl_among_passing() {
2238 let base = vec![0.10, 0.00, 0.20];
2239 let evals = vec![
2240 (25usize, 21.0, vec![0.10, 0.05, 0.20]), (50, 18.0, vec![0.40, 0.00, 0.10]), (75, 19.0, vec![0.15, 0.05, 0.25]), (100, 18.5, vec![0.20, 0.30, 0.20]), ];
2245 let sel = select_checkpoint(&evals, &base, 0.35, 0.10);
2246 assert_eq!(sel, Some(2), "step 75 is the lowest-ppl PASSING checkpoint");
2247 }
2248
2249 #[test]
2252 fn gate_all_fail_is_identity() {
2253 let base = vec![0.0, 0.0, 0.0];
2254 let evals = vec![
2255 (25usize, 15.0, vec![0.50, 0.0, 0.0]),
2256 (50, 14.0, vec![0.0, 0.36, 0.0]),
2257 (75, 13.0, vec![0.0, 0.0, 0.11]), ];
2259 assert_eq!(select_checkpoint(&evals, &base, 0.35, 0.10), None);
2260 }
2261
2262 #[test]
2265 fn gate_boundaries_and_tie_break() {
2266 let base = vec![0.25];
2267 assert!(gate_pass(&[0.35], &base, 0.35, 0.10), "== threshold passes");
2268 assert!(
2269 gate_pass(&[0.35], &[0.25], 0.35, 0.10),
2270 "== base+slack passes"
2271 );
2272 assert!(!gate_pass(&[0.351], &base, 0.35, 0.10));
2273 assert!(!gate_pass(&[0.30], &[0.10], 0.35, 0.10), "0.30 > 0.10+0.10");
2274 let evals = vec![(25usize, 20.0, vec![0.10]), (50, 20.0, vec![0.10])];
2275 assert_eq!(
2276 select_checkpoint(&evals, &base, 0.35, 0.10),
2277 Some(0),
2278 "equal ppl → earliest checkpoint"
2279 );
2280 }
2281}