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 { steps: 300, lr: 5e-5, kl_w: 0.7, eval_every: 25, bs: 2, seq: 512, seed: 0 }
58 }
59}
60
61#[derive(Clone, Debug)]
64pub struct FcdReport {
65 pub converted: Vec<usize>,
66 pub teacher_ppl: f64,
68 pub ppl_start: f64,
70 pub ppl_best: f64,
72 pub best_step: usize,
73 pub ppl_final: f64,
75 pub steps_run: usize,
76 pub sec_per_step: f64,
77 pub losses: Vec<(f64, f64)>,
79 pub gate: Option<GateReport>,
81}
82
83#[derive(Clone, Debug)]
85pub struct GateReport {
86 pub baseline: Vec<f64>,
88 pub evals: Vec<(usize, f64, Vec<f64>, bool)>,
90 pub chosen: Option<usize>,
93}
94
95pub fn loop_score(ids: &[u32]) -> f64 {
100 if ids.len() < 5 {
101 return 0.0;
102 }
103 let grams: std::collections::HashSet<&[u32]> = ids.windows(4).collect();
104 1.0 - grams.len() as f64 / ids.windows(4).count() as f64
105}
106
107#[derive(Clone, Debug)]
111pub struct GenGateCfg {
112 pub prompts: Vec<Vec<u32>>,
115 pub gen_tokens: usize,
116 pub threshold: f64,
118 pub baseline_slack: f64,
120}
121
122impl GenGateCfg {
123 pub fn standard(va: &[u32]) -> Option<Self> {
126 let l = va.len().saturating_sub(500);
127 if l < 400 {
128 return None;
129 }
130 let prompts = [l / 10, l / 2, 8 * l / 10]
131 .iter()
132 .map(|&off| va[off..off + 400].to_vec())
133 .collect();
134 Some(Self { prompts, gen_tokens: 60, threshold: 0.35, baseline_slack: 0.10 })
135 }
136}
137
138pub fn gate_pass(scores: &[f64], baseline: &[f64], threshold: f64, slack: f64) -> bool {
142 scores.iter().zip(baseline).all(|(&s, &b)| s <= threshold && s <= b + slack)
143}
144
145pub fn select_checkpoint(
150 evals: &[(usize, f64, Vec<f64>)],
151 baseline: &[f64],
152 threshold: f64,
153 slack: f64,
154) -> Option<usize> {
155 let mut best: Option<usize> = None;
156 for (i, (_, ppl, scores)) in evals.iter().enumerate() {
157 if !gate_pass(scores, baseline, threshold, slack) {
158 continue;
159 }
160 if best.map(|b| *ppl < evals[b].1).unwrap_or(true) {
161 best = Some(i);
162 }
163 }
164 best
165}
166
167enum FcdAttn {
172 Full {
173 wq: Vec<f32>,
174 wk: Vec<f32>,
175 wv: Vec<f32>,
176 wo: Vec<f32>,
177 q_norm: Option<Vec<f32>>,
178 k_norm: Option<Vec<f32>>,
179 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
180 output_gate: bool,
183 },
184 Gdn {
187 wqkv: Vec<f32>,
188 wz: Vec<f32>,
189 wa: Vec<f32>,
190 wb: Vec<f32>,
191 conv: Vec<f32>,
192 a_log: Vec<f32>,
193 dt_bias: Vec<f32>,
194 norm: Vec<f32>,
195 wout: Vec<f32>,
196 },
197}
198
199struct FcdLayer {
200 attn: FcdAttn,
201 inter: usize,
202 iln: Vec<f32>,
204 pln: Vec<f32>,
205 gate: Vec<f32>,
206 up: Vec<f32>,
207 down: Vec<f32>,
208}
209
210#[derive(Clone, Copy)]
212struct GdnDims {
213 nv: usize,
214 nk: usize,
215 dk: usize,
216 dv: usize,
217 kk: usize,
218}
219
220impl GdnDims {
221 fn c_dim(&self) -> usize {
222 2 * self.nk * self.dk + self.nv * self.dv
223 }
224 fn vd(&self) -> usize {
225 self.nv * self.dv
226 }
227}
228
229pub struct FcdModel {
231 pub hidden: usize,
232 pub nh: usize,
233 pub nkv: usize,
234 pub hd: usize,
235 pub nl: usize,
236 pub vocab: usize,
237 eps: f64,
238 gemma: bool,
239 rotary_dim: usize,
240 inv_freq: Vec<f64>,
241 embed: Vec<f32>,
243 lm_head: Option<Vec<f32>>,
244 final_norm: Vec<f32>,
245 layers: Vec<FcdLayer>,
246 o1_flags: Vec<bool>,
248 nys: NysCfg,
249 gdn: Option<GdnDims>,
251 pool: Option<Arc<Pool>>,
252}
253
254fn deq(model: &CmfModel, name: &str) -> Result<Vec<f32>, String> {
255 let e = model
256 .tensor(name)
257 .ok_or_else(|| format!("tensor '{name}' not found"))?;
258 let mut out = vec![0f32; e.n_elems()];
259 cortiq_core::quant::dequant_tensor(e, model.entry_bytes(e), &mut out)?;
260 Ok(out)
261}
262
263impl FcdModel {
264 pub fn from_cmf(model: &CmfModel, o1: &O1Cfg) -> Result<Self, String> {
267 let arch = model.arch().clone();
268 let has_linear = arch
269 .layer_types
270 .iter()
271 .any(|t| matches!(t, LayerType::LinearAttention));
272 let gdn = if has_linear {
273 let lc = arch.linear_core.as_ref().ok_or_else(|| {
274 "model has linear layers but no arch.linear_core".to_string()
275 })?;
276 if lc.kind != "gated_delta_net" {
277 return Err(format!(
278 "linear core '{}' has no FCD backward (only gated_delta_net)",
279 lc.kind
280 ));
281 }
282 Some(GdnDims {
283 nv: lc.num_heads,
284 nk: arch
285 .linear_num_key_heads
286 .ok_or("linear core needs arch.linear_num_key_heads")?,
287 dk: arch
288 .linear_key_head_dim
289 .ok_or("linear core needs arch.linear_key_head_dim")?,
290 dv: lc.value_head_dim,
291 kk: arch
292 .linear_conv_kernel_dim
293 .ok_or("linear core needs arch.linear_conv_kernel_dim")?,
294 })
295 } else {
296 None
297 };
298 let (nh, nkv, hd, h) = (
299 arch.num_attention_heads,
300 arch.num_kv_heads,
301 arch.head_dim,
302 arch.hidden_size,
303 );
304 let embed = deq(model, "model.embed_tokens.weight")?;
305 let lm_head = if model.tensor("lm_head.weight").is_some() {
306 Some(deq(model, "lm_head.weight")?)
307 } else if arch.tie_word_embeddings {
308 None
309 } else {
310 return Err("no lm_head.weight and tie_word_embeddings is false".into());
311 };
312 let final_norm = deq(model, "model.norm.weight")?;
313
314 let mut layers = Vec::with_capacity(arch.num_layers);
315 for li in 0..arch.num_layers {
316 let p = format!("model.layers.{li}.");
317 if model.tensor(&format!("{p}mlp.gate.weight")).is_some() {
318 return Err(format!(
319 "layer {li} is MoE — FCD polish supports dense FFN only"
320 ));
321 }
322 let attn = match arch.layer_types.get(li) {
323 Some(LayerType::LinearAttention) => {
324 let la = |n: &str| deq(model, &format!("{p}linear_attn.{n}"));
325 FcdAttn::Gdn {
326 wqkv: la("in_proj_qkv.weight")?,
327 wz: la("in_proj_z.weight")?,
328 wa: la("in_proj_a.weight")?,
329 wb: la("in_proj_b.weight")?,
330 conv: la("conv1d.weight")?,
331 a_log: la("A_log")?,
332 dt_bias: la("dt_bias")?,
333 norm: la("norm.weight")?,
334 wout: la("out_proj.weight")?,
335 }
336 }
337 _ => {
338 let wq = deq(model, &format!("{p}self_attn.q_proj.weight"))?;
339 let output_gate = wq.len() == 2 * nh * hd * h;
340 let opt = |n: &str| -> Option<Vec<f32>> {
341 model
342 .tensor(&format!("{p}self_attn.{n}"))
343 .and_then(|_| deq(model, &format!("{p}self_attn.{n}")).ok())
344 };
345 let bias = match (
346 opt("q_proj.bias"),
347 opt("k_proj.bias"),
348 opt("v_proj.bias"),
349 ) {
350 (Some(a), Some(b), Some(c)) => Some((a, b, c)),
351 _ => None,
352 };
353 FcdAttn::Full {
354 wq,
355 wk: deq(model, &format!("{p}self_attn.k_proj.weight"))?,
356 wv: deq(model, &format!("{p}self_attn.v_proj.weight"))?,
357 wo: deq(model, &format!("{p}self_attn.o_proj.weight"))?,
358 q_norm: opt("q_norm.weight"),
359 k_norm: opt("k_norm.weight"),
360 bias,
361 output_gate,
362 }
363 }
364 };
365 let gate = deq(model, &format!("{p}mlp.gate_proj.weight"))?;
366 let inter = gate.len() / h;
367 layers.push(FcdLayer {
368 attn,
369 inter,
370 iln: deq(model, &format!("{p}input_layernorm.weight"))?,
371 pln: deq(model, &format!("{p}post_attention_layernorm.weight"))?,
372 gate,
373 up: deq(model, &format!("{p}mlp.up_proj.weight"))?,
374 down: deq(model, &format!("{p}mlp.down_proj.weight"))?,
375 });
376 }
377
378 let rotary_dim = ((hd as f32 * arch.partial_rotary_factor) as usize).max(2).min(hd);
379 let base = arch.rope_theta;
380 let inv_freq: Vec<f64> = (0..rotary_dim / 2)
381 .map(|i| 1.0 / base.powf(2.0 * i as f64 / rotary_dim as f64))
382 .collect();
383 let mut flags = o1.layer_flags(arch.num_layers);
384 flags.resize(arch.num_layers, false);
385 for (li, f) in flags.iter_mut().enumerate() {
388 if *f && !matches!(layers[li].attn, FcdAttn::Full { .. }) {
389 *f = false;
390 }
391 }
392 Ok(Self {
393 hidden: h,
394 nh,
395 nkv,
396 hd,
397 nl: arch.num_layers,
398 vocab: arch.vocab_size.min(embed.len() / h),
399 eps: arch.rms_norm_eps,
400 gemma: matches!(arch.norm_style, NormStyle::Gemma),
401 rotary_dim,
402 inv_freq,
403 embed,
404 lm_head,
405 final_norm,
406 layers,
407 o1_flags: flags,
408 nys: NysCfg { m: o1.m, w: o1.w, sink: o1.sink, prefill: None },
411 gdn,
412 pool: Pool::from_env(),
413 })
414 }
415
416 pub fn converted(&self) -> Vec<usize> {
418 (0..self.nl).filter(|&i| self.o1_flags[i]).collect()
419 }
420
421 fn head_weight(&self) -> &[f32] {
422 self.lm_head.as_deref().unwrap_or(&self.embed)
423 }
424}
425
426const PARAMS_PER_LAYER: usize = 5; pub struct TrainState {
433 pub layers: Vec<usize>,
434 pub data: Vec<Vec<f32>>,
436 grad: Vec<Vec<f32>>,
437 m1: Vec<Vec<f32>>,
438 m2: Vec<Vec<f32>>,
439 step_t: u64,
440}
441
442impl TrainState {
443 pub fn new(fm: &FcdModel) -> Self {
444 let layers = fm.converted();
445 let mut data = Vec::with_capacity(layers.len() * PARAMS_PER_LAYER);
446 for &li in &layers {
447 let l = &fm.layers[li];
448 data.push(l.iln.clone());
449 data.push(l.pln.clone());
450 data.push(l.gate.clone());
451 data.push(l.up.clone());
452 data.push(l.down.clone());
453 }
454 let zeros: Vec<Vec<f32>> = data.iter().map(|d| vec![0f32; d.len()]).collect();
455 Self {
456 layers,
457 grad: zeros.clone(),
458 m1: zeros.clone(),
459 m2: zeros,
460 data,
461 step_t: 0,
462 }
463 }
464
465 fn slot(&self, li: usize) -> Option<usize> {
466 self.layers.iter().position(|&x| x == li)
467 }
468
469 #[doc(hidden)]
471 pub fn grads(&self) -> &[Vec<f32>] {
472 &self.grad
473 }
474
475 fn zero_grad(&mut self) {
476 for g in &mut self.grad {
477 for v in g.iter_mut() {
478 *v = 0.0;
479 }
480 }
481 }
482
483 fn clip_and_step(&mut self, lr: f64) -> f64 {
486 let mut sq = 0f64;
487 for g in &self.grad {
488 for &v in g {
489 sq += (v as f64) * (v as f64);
490 }
491 }
492 let gn = sq.sqrt();
493 let scale = if gn > 1.0 { 1.0 / (gn + 1e-6) } else { 1.0 };
494 self.step_t += 1;
495 let bc1 = 1.0 - ADAM_B1.powi(self.step_t as i32);
496 let bc2 = 1.0 - ADAM_B2.powi(self.step_t as i32);
497 for p in 0..self.data.len() {
498 let (d, g, m, v) = (
499 &mut self.data[p],
500 &self.grad[p],
501 &mut self.m1[p],
502 &mut self.m2[p],
503 );
504 for i in 0..d.len() {
505 let gi = g[i] as f64 * scale;
506 let mi = ADAM_B1 * m[i] as f64 + (1.0 - ADAM_B1) * gi;
507 let vi = ADAM_B2 * v[i] as f64 + (1.0 - ADAM_B2) * gi * gi;
508 m[i] = mi as f32;
509 v[i] = vi as f32;
510 let upd = (mi / bc1) / ((vi / bc2).sqrt() + ADAM_EPS) + ADAM_WD * d[i] as f64;
511 d[i] = (d[i] as f64 - lr * upd) as f32;
512 }
513 }
514 gn
515 }
516}
517
518#[derive(Clone, Copy)]
521struct LnFfn<'a> {
522 iln: &'a [f32],
523 pln: &'a [f32],
524 gate: &'a [f32],
525 up: &'a [f32],
526 down: &'a [f32],
527}
528
529fn ln_ffn<'a>(fm: &'a FcdModel, ts: Option<&'a TrainState>, li: usize) -> LnFfn<'a> {
530 if let Some(t) = ts {
531 if let Some(s) = t.slot(li) {
532 let b = s * PARAMS_PER_LAYER;
533 return LnFfn {
534 iln: &t.data[b],
535 pln: &t.data[b + 1],
536 gate: &t.data[b + 2],
537 up: &t.data[b + 3],
538 down: &t.data[b + 4],
539 };
540 }
541 }
542 let l = &fm.layers[li];
543 LnFfn { iln: &l.iln, pln: &l.pln, gate: &l.gate, up: &l.up, down: &l.down }
544}
545
546enum AttnActs {
550 Full {
551 qpre: Vec<f32>,
552 kpre: Vec<f32>,
553 vproj: Vec<f32>,
554 qrot: Vec<f32>,
555 krot: Vec<f32>,
556 qinv: Vec<f32>,
557 kinv: Vec<f32>,
558 ao: Vec<f32>,
561 gate_pre: Vec<f32>,
563 },
564 Gdn { qkv: Vec<f32>, z: Vec<f32>, a: Vec<f32>, b: Vec<f32> },
567}
568
569struct LayerActs {
570 inv1: Vec<f32>,
571 attn: AttnActs,
572 h1: Vec<f32>,
573 n2: Vec<f32>,
574 inv2: Vec<f32>,
575 gpre: Vec<f32>,
576 upre: Vec<f32>,
577 act: Vec<f32>,
578}
579
580struct SendMut<T>(*mut T);
582unsafe impl<T> Send for SendMut<T> {}
583unsafe impl<T> Sync for SendMut<T> {}
584impl<T> SendMut<T> {
585 #[inline]
586 unsafe fn at(&self, i: usize) -> *mut T {
587 unsafe { self.0.add(i) }
588 }
589}
590
591impl FcdModel {
592 fn qk_norm_rope(
596 &self,
597 x: &mut [f32],
598 norm: Option<&[f32]>,
599 heads: usize,
600 t: usize,
601 inv_out: &mut [f32],
602 ) {
603 let hd = self.hd;
604 let n = x.len() / (heads * hd);
605 for r in 0..n {
606 let pos = r % t;
607 for hh in 0..heads {
608 let s = (r * heads + hh) * hd;
609 let head = &mut x[s..s + hd];
610 if let Some(w) = norm {
611 let mut inv = [0f32; 1];
612 let mut y = [0f32; 256];
613 debug_assert!(hd <= 256);
614 ops::rmsnorm_fwd(head, w, self.eps, self.gemma, &mut y[..hd], &mut inv);
615 head.copy_from_slice(&y[..hd]);
616 inv_out[r * heads + hh] = inv[0];
617 }
618 ops::rope_fwd(&mut head[..self.rotary_dim], pos, &self.inv_freq);
619 }
620 }
621 }
622
623 #[allow(clippy::too_many_arguments)]
629 fn layer_forward(
630 &self,
631 li: usize,
632 h_in: &[f32],
633 b: usize,
634 t: usize,
635 wts: &LnFfn,
636 nystrom: bool,
637 want_acts: bool,
638 ) -> (Vec<f32>, Option<LayerActs>) {
639 let hsz = self.hidden;
640 let n = b * t;
641 let l = &self.layers[li];
642 let pool = self.pool.as_deref();
643
644 let mut n1 = vec![0f32; n * hsz];
645 let mut inv1 = vec![0f32; n];
646 ops::rmsnorm_fwd(h_in, wts.iln, self.eps, self.gemma, &mut n1, &mut inv1);
647
648 let (attn_out, attn_acts) = match &l.attn {
649 FcdAttn::Full { .. } => self.full_attn_fwd(&l.attn, &n1, b, t, nystrom),
650 FcdAttn::Gdn { .. } => self.gdn_attn_fwd(&l.attn, &n1, b, t),
651 };
652
653 let mut h1 = h_in.to_vec();
654 for (a, &x) in h1.iter_mut().zip(&attn_out) {
655 *a += x;
656 }
657
658 let mut n2 = vec![0f32; n * hsz];
659 let mut inv2 = vec![0f32; n];
660 ops::rmsnorm_fwd(&h1, wts.pln, self.eps, self.gemma, &mut n2, &mut inv2);
661
662 let inter = l.inter;
663 let mut gpre = vec![0f32; n * inter];
664 ops::gemm_nt(&n2, wts.gate, &mut gpre, n, hsz, inter, pool);
665 let mut upre = vec![0f32; n * inter];
666 ops::gemm_nt(&n2, wts.up, &mut upre, n, hsz, inter, pool);
667 let mut act = vec![0f32; n * inter];
668 for i in 0..n * inter {
669 act[i] = ops::silu(gpre[i]) * upre[i];
670 }
671 let mut ffn = vec![0f32; n * hsz];
672 ops::gemm_nt(&act, wts.down, &mut ffn, n, inter, hsz, pool);
673 let mut h2 = h1.clone();
674 for (a, &x) in h2.iter_mut().zip(&ffn) {
675 *a += x;
676 }
677
678 let acts = want_acts.then_some(LayerActs {
679 inv1,
680 attn: attn_acts,
681 h1,
682 n2,
683 inv2,
684 gpre,
685 upre,
686 act,
687 });
688 (h2, acts)
689 }
690
691 fn full_attn_fwd(
695 &self,
696 attn: &FcdAttn,
697 n1: &[f32],
698 b: usize,
699 t: usize,
700 nystrom: bool,
701 ) -> (Vec<f32>, AttnActs) {
702 let FcdAttn::Full { wq, wk, wv, wo, q_norm, k_norm, bias, output_gate } = attn else {
703 unreachable!("full_attn_fwd on a non-Full layer");
704 };
705 let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
706 let n = b * t;
707 let pool = self.pool.as_deref();
708 let qdim = nh * hd;
709 let kvdim = nkv * hd;
710 let rep = nh / nkv;
711 let qrows = if *output_gate { 2 * qdim } else { qdim };
712
713 let mut qraw = vec![0f32; n * qrows];
714 ops::gemm_nt(n1, wq, &mut qraw, n, hsz, qrows, pool);
715 let mut kpre = vec![0f32; n * kvdim];
716 ops::gemm_nt(n1, wk, &mut kpre, n, hsz, kvdim, pool);
717 let mut vproj = vec![0f32; n * kvdim];
718 ops::gemm_nt(n1, wv, &mut vproj, n, hsz, kvdim, pool);
719 if let Some((bq, bk, bv)) = bias {
720 for r in 0..n {
721 for (x, bb) in qraw[r * qrows..(r + 1) * qrows].iter_mut().zip(bq) {
722 *x += bb;
723 }
724 for (x, bb) in kpre[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bk) {
725 *x += bb;
726 }
727 for (x, bb) in vproj[r * kvdim..(r + 1) * kvdim].iter_mut().zip(bv) {
728 *x += bb;
729 }
730 }
731 }
732 let (qpre, gate_pre) = if *output_gate {
734 let mut qh = vec![0f32; n * qdim];
735 let mut gp = vec![0f32; n * qdim];
736 for r in 0..n {
737 for h in 0..nh {
738 let src = r * qrows + 2 * h * hd;
739 let dst = r * qdim + h * hd;
740 qh[dst..dst + hd].copy_from_slice(&qraw[src..src + hd]);
741 gp[dst..dst + hd].copy_from_slice(&qraw[src + hd..src + 2 * hd]);
742 }
743 }
744 (qh, gp)
745 } else {
746 (qraw, Vec::new())
747 };
748
749 let mut qrot = qpre.clone();
750 let mut krot = kpre.clone();
751 let mut qinv = vec![0f32; n * nh];
752 let mut kinv = vec![0f32; n * nkv];
753 self.qk_norm_rope(&mut qrot, q_norm.as_deref(), nh, t, &mut qinv);
754 self.qk_norm_rope(&mut krot, k_norm.as_deref(), nkv, t, &mut kinv);
755
756 let mut ao = vec![0f32; n * qdim];
758 {
759 let units = b * nh;
760 let aop = SendMut(ao.as_mut_ptr());
761 let qr = &qrot;
762 let kr = &krot;
763 let vr = &vproj;
764 let nys = self.nys;
765 let run_unit = |u: usize| {
766 let (bi, h) = (u / nh, u % nh);
767 let g = h / rep;
768 if nystrom {
769 let mut q64 = vec![0f64; t * hd];
771 let mut k64 = vec![0f64; t * hd];
772 let mut v64 = vec![0f64; t * hd];
773 for p in 0..t {
774 let r = bi * t + p;
775 for c in 0..hd {
776 q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
777 k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
778 v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
779 }
780 }
781 let mut o64 = vec![0f64; t * hd];
782 ops::nystrom_head_fwd(&q64, &k64, &v64, t, hd, hd, &nys, &mut o64);
783 for p in 0..t {
784 let r = bi * t + p;
785 for c in 0..hd {
786 unsafe {
788 *aop.at(r * qdim + h * hd + c) = o64[p * hd + c] as f32;
789 }
790 }
791 }
792 } else {
793 let mut q32 = vec![0f32; t * hd];
794 let mut k32 = vec![0f32; t * hd];
795 let mut v32 = vec![0f32; t * hd];
796 for p in 0..t {
797 let r = bi * t + p;
798 q32[p * hd..(p + 1) * hd]
799 .copy_from_slice(&qr[r * qdim + h * hd..r * qdim + (h + 1) * hd]);
800 k32[p * hd..(p + 1) * hd]
801 .copy_from_slice(&kr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
802 v32[p * hd..(p + 1) * hd]
803 .copy_from_slice(&vr[r * kvdim + g * hd..r * kvdim + (g + 1) * hd]);
804 }
805 let mut o32 = vec![0f32; t * hd];
806 ops::attn_head_fwd(&q32, &k32, &v32, t, hd, hd, &mut o32);
807 for p in 0..t {
808 let r = bi * t + p;
809 for c in 0..hd {
810 unsafe {
812 *aop.at(r * qdim + h * hd + c) = o32[p * hd + c];
813 }
814 }
815 }
816 }
817 };
818 match pool {
819 Some(p) if units > 1 => p.run(&|widx, nw| {
820 for u in (widx..units).step_by(nw) {
821 run_unit(u);
822 }
823 }),
824 _ => {
825 for u in 0..units {
826 run_unit(u);
827 }
828 }
829 }
830 }
831
832 let ao_eff: Vec<f32> = if *output_gate {
834 ao.iter()
835 .zip(&gate_pre)
836 .map(|(&a, &g)| a * (1.0 / (1.0 + (-g).exp())))
837 .collect()
838 } else {
839 ao.clone()
840 };
841 let mut attn_out = vec![0f32; n * hsz];
842 ops::gemm_nt(&ao_eff, wo, &mut attn_out, n, qdim, hsz, pool);
843 (
844 attn_out,
845 AttnActs::Full { qpre, kpre, vproj, qrot, krot, qinv, kinv, ao, gate_pre },
846 )
847 }
848
849 fn gdn_attn_fwd(
854 &self,
855 attn: &FcdAttn,
856 n1: &[f32],
857 b: usize,
858 t: usize,
859 ) -> (Vec<f32>, AttnActs) {
860 let FcdAttn::Gdn { wqkv, wz, wa, wb, conv, a_log, dt_bias, norm, wout } = attn else {
861 unreachable!("gdn_attn_fwd on a non-GDN layer");
862 };
863 let d = self.gdn.expect("gdn layer without gdn dims");
864 let (hsz, n) = (self.hidden, b * t);
865 let pool = self.pool.as_deref();
866 let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
867
868 let mut qkv = vec![0f32; n * c_dim];
869 ops::gemm_nt(n1, wqkv, &mut qkv, n, hsz, c_dim, pool);
870 let mut z = vec![0f32; n * vd];
871 ops::gemm_nt(n1, wz, &mut z, n, hsz, vd, pool);
872 let mut a = vec![0f32; n * nv];
873 ops::gemm_nt(n1, wa, &mut a, n, hsz, nv, pool);
874 let mut bstr = vec![0f32; n * nv];
875 ops::gemm_nt(n1, wb, &mut bstr, n, hsz, nv, pool);
876
877 let cfg = ops::GdnSeqCfg {
878 nv: d.nv,
879 nk: d.nk,
880 dk: d.dk,
881 dv: d.dv,
882 kk: d.kk,
883 rms_eps: self.eps,
884 conv,
885 a_log,
886 dt_bias,
887 norm,
888 };
889 let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
891 let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
892 let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
893 let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
894 let mut pre64 = vec![0f64; n * c_dim];
895 let mut cq64 = vec![0f64; n * c_dim];
896 for bi in 0..b {
897 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
898 ops::gdn_conv_fwd(
899 &qkv64[r.clone()],
900 t,
901 c_dim,
902 d.kk,
903 conv,
904 &mut pre64[r.clone()],
905 &mut cq64[r],
906 );
907 }
908 let mut of = vec![0f32; n * vd];
909 {
910 let units = b * d.nk;
911 let rep_v = d.nv / d.nk;
912 let ofp = SendMut(of.as_mut_ptr());
913 let (cqr, zr, ar, br) = (&cq64, &z64, &a64, &b64);
914 let cfg_ref = &cfg;
915 let run_unit = |u: usize| {
916 let (bi, ko) = (u / d.nk, u % d.nk);
917 let mut local = vec![0f64; t * vd];
918 ops::gdn_group_fwd(
919 &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
920 &zr[bi * t * vd..(bi + 1) * t * vd],
921 &ar[bi * t * nv..(bi + 1) * t * nv],
922 &br[bi * t * nv..(bi + 1) * t * nv],
923 t,
924 cfg_ref,
925 ko,
926 &mut local,
927 );
928 for hh in 0..rep_v {
929 let h = ko * rep_v + hh;
930 for p in 0..t {
931 for dj in 0..d.dv {
932 unsafe {
934 *ofp.at((bi * t + p) * vd + h * d.dv + dj) =
935 local[p * vd + h * d.dv + dj] as f32;
936 }
937 }
938 }
939 }
940 };
941 match pool {
942 Some(p) if units > 1 => p.run(&|widx, nw| {
943 for u in (widx..units).step_by(nw) {
944 run_unit(u);
945 }
946 }),
947 _ => {
948 for u in 0..units {
949 run_unit(u);
950 }
951 }
952 }
953 }
954 let mut attn_out = vec![0f32; n * hsz];
955 ops::gemm_nt(&of, wout, &mut attn_out, n, vd, hsz, pool);
956 (attn_out, AttnActs::Gdn { qkv, z, a, b: bstr })
957 }
958
959 #[allow(clippy::too_many_arguments)]
963 fn layer_backward(
964 &self,
965 li: usize,
966 h_in: &[f32],
967 b: usize,
968 t: usize,
969 wts: &LnFfn,
970 nystrom: bool,
971 acts: &LayerActs,
972 dh2: &[f32],
973 mut grads: Option<&mut [Vec<f32>]>,
974 ) -> Vec<f32> {
975 let hsz = self.hidden;
976 let n = b * t;
977 let l = &self.layers[li];
978 let pool = self.pool.as_deref();
979 let inter = l.inter;
980
981 let mut dact = vec![0f32; n * inter];
983 ops::gemm_dx(dh2, wts.down, &mut dact, n, inter, hsz, pool);
984 if let Some(g) = grads.as_deref_mut() {
985 ops::gemm_dw(dh2, &acts.act, &mut g[4], n, inter, hsz, pool);
986 }
987 let mut dg = vec![0f32; n * inter];
988 let mut du = vec![0f32; n * inter];
989 for i in 0..n * inter {
990 dg[i] = dact[i] * acts.upre[i] * ops::silu_bwd(acts.gpre[i]);
991 du[i] = dact[i] * ops::silu(acts.gpre[i]);
992 }
993 let mut dn2 = vec![0f32; n * hsz];
994 ops::gemm_dx(&dg, wts.gate, &mut dn2, n, hsz, inter, pool);
995 ops::gemm_dx(&du, wts.up, &mut dn2, n, hsz, inter, pool);
996 if let Some(g) = grads.as_deref_mut() {
997 ops::gemm_dw(&dg, &acts.n2, &mut g[2], n, hsz, inter, pool);
998 ops::gemm_dw(&du, &acts.n2, &mut g[3], n, hsz, inter, pool);
999 }
1000
1001 let mut dh1 = dh2.to_vec();
1002 ops::rmsnorm_bwd(
1003 &acts.h1,
1004 wts.pln,
1005 &acts.inv2,
1006 &dn2,
1007 self.gemma,
1008 &mut dh1,
1009 grads.as_deref_mut().map(|g| &mut g[1][..]),
1010 );
1011
1012 let dn1 = match &l.attn {
1014 FcdAttn::Full { .. } => {
1015 self.full_attn_bwd(&l.attn, &acts.attn, &dh1, b, t, nystrom)
1016 }
1017 FcdAttn::Gdn { .. } => self.gdn_attn_bwd(&l.attn, &acts.attn, &dh1, b, t),
1018 };
1019
1020 let mut dh_in = dh1.clone();
1021 ops::rmsnorm_bwd(
1022 h_in,
1023 wts.iln,
1024 &acts.inv1,
1025 &dn1,
1026 self.gemma,
1027 &mut dh_in,
1028 grads.map(|g| &mut g[0][..]),
1029 );
1030 dh_in
1031 }
1032
1033 fn full_attn_bwd(
1037 &self,
1038 attn: &FcdAttn,
1039 acts: &AttnActs,
1040 dattn: &[f32],
1041 b: usize,
1042 t: usize,
1043 nystrom: bool,
1044 ) -> Vec<f32> {
1045 let FcdAttn::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, .. } = attn else {
1046 unreachable!("full_attn_bwd on a non-Full layer");
1047 };
1048 let AttnActs::Full { qpre, kpre, vproj, qrot, krot, qinv, kinv, ao, gate_pre } = acts
1049 else {
1050 unreachable!("acts mismatch");
1051 };
1052 let (hsz, nh, nkv, hd) = (self.hidden, self.nh, self.nkv, self.hd);
1053 let n = b * t;
1054 let pool = self.pool.as_deref();
1055 let qdim = nh * hd;
1056 let kvdim = nkv * hd;
1057 let rep = nh / nkv;
1058 let qrows = if *output_gate { 2 * qdim } else { qdim };
1059
1060 let mut dao_eff = vec![0f32; n * qdim];
1061 ops::gemm_dx(dattn, wo, &mut dao_eff, n, qdim, hsz, pool);
1062 let (dao, dgate) = if *output_gate {
1064 let mut dao = vec![0f32; n * qdim];
1065 let mut dgp = vec![0f32; n * qdim];
1066 for i in 0..n * qdim {
1067 let sig = 1.0 / (1.0 + (-gate_pre[i]).exp());
1068 dao[i] = dao_eff[i] * sig;
1069 dgp[i] = dao_eff[i] * ao[i] * sig * (1.0 - sig);
1070 }
1071 (dao, dgp)
1072 } else {
1073 (dao_eff, Vec::new())
1074 };
1075
1076 let mut dqrot = vec![0f32; n * qdim];
1077 let mut dkrot = vec![0f32; n * kvdim];
1078 let mut dvproj = vec![0f32; n * kvdim];
1079 {
1080 let units = b * nkv;
1083 let dqp = SendMut(dqrot.as_mut_ptr());
1084 let dkp = SendMut(dkrot.as_mut_ptr());
1085 let dvp = SendMut(dvproj.as_mut_ptr());
1086 let (qr, kr, vr) = (qrot, krot, vproj);
1087 let daor = &dao;
1088 let nys = self.nys;
1089 let run_unit = |u: usize| {
1090 let (bi, g) = (u / nkv, u % nkv);
1091 let mut k64 = vec![0f64; t * hd];
1092 let mut v64 = vec![0f64; t * hd];
1093 for p in 0..t {
1094 let r = bi * t + p;
1095 for c in 0..hd {
1096 k64[p * hd + c] = kr[r * kvdim + g * hd + c] as f64;
1097 v64[p * hd + c] = vr[r * kvdim + g * hd + c] as f64;
1098 }
1099 }
1100 let mut dk64 = vec![0f64; t * hd];
1101 let mut dv64 = vec![0f64; t * hd];
1102 let mut q64 = vec![0f64; t * hd];
1103 let mut do64 = vec![0f64; t * hd];
1104 let mut dq64 = vec![0f64; t * hd];
1105 for hh in 0..rep {
1106 let h = g * rep + hh;
1107 for p in 0..t {
1108 let r = bi * t + p;
1109 for c in 0..hd {
1110 q64[p * hd + c] = qr[r * qdim + h * hd + c] as f64;
1111 do64[p * hd + c] = daor[r * qdim + h * hd + c] as f64;
1112 }
1113 }
1114 for v in dq64.iter_mut() {
1115 *v = 0.0;
1116 }
1117 if nystrom {
1118 ops::nystrom_head_bwd(
1119 &q64, &k64, &v64, &do64, t, hd, hd, &nys, &mut dq64, &mut dk64,
1120 &mut dv64,
1121 );
1122 } else {
1123 ops::attn_head_bwd(
1124 &q64, &k64, &v64, &do64, t, hd, hd, &mut dq64, &mut dk64, &mut dv64,
1125 );
1126 }
1127 for p in 0..t {
1128 let r = bi * t + p;
1129 for c in 0..hd {
1130 unsafe {
1132 *dqp.at(r * qdim + h * hd + c) = dq64[p * hd + c] as f32;
1133 }
1134 }
1135 }
1136 }
1137 for p in 0..t {
1138 let r = bi * t + p;
1139 for c in 0..hd {
1140 unsafe {
1142 *dkp.at(r * kvdim + g * hd + c) = dk64[p * hd + c] as f32;
1143 *dvp.at(r * kvdim + g * hd + c) = dv64[p * hd + c] as f32;
1144 }
1145 }
1146 }
1147 };
1148 match pool {
1149 Some(p) if units > 1 => p.run(&|widx, nw| {
1150 for u in (widx..units).step_by(nw) {
1151 run_unit(u);
1152 }
1153 }),
1154 _ => {
1155 for u in 0..units {
1156 run_unit(u);
1157 }
1158 }
1159 }
1160 }
1161
1162 let mut dqpre = vec![0f32; n * qdim];
1164 let mut dkpre = vec![0f32; n * kvdim];
1165 for r in 0..n {
1166 let pos = r % t;
1167 for h in 0..nh {
1168 let s = r * qdim + h * hd;
1169 ops::rope_bwd(&mut dqrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1170 match q_norm {
1171 Some(w) => ops::rmsnorm_bwd(
1172 &qpre[s..s + hd],
1173 w,
1174 &qinv[r * nh + h..r * nh + h + 1],
1175 &dqrot[s..s + hd],
1176 self.gemma,
1177 &mut dqpre[s..s + hd],
1178 None,
1179 ),
1180 None => dqpre[s..s + hd].copy_from_slice(&dqrot[s..s + hd]),
1181 }
1182 }
1183 for g in 0..nkv {
1184 let s = r * kvdim + g * hd;
1185 ops::rope_bwd(&mut dkrot[s..s + self.rotary_dim], pos, &self.inv_freq);
1186 match k_norm {
1187 Some(w) => ops::rmsnorm_bwd(
1188 &kpre[s..s + hd],
1189 w,
1190 &kinv[r * nkv + g..r * nkv + g + 1],
1191 &dkrot[s..s + hd],
1192 self.gemma,
1193 &mut dkpre[s..s + hd],
1194 None,
1195 ),
1196 None => dkpre[s..s + hd].copy_from_slice(&dkrot[s..s + hd]),
1197 }
1198 }
1199 }
1200
1201 let dqraw: Vec<f32> = if *output_gate {
1203 let mut dq = vec![0f32; n * qrows];
1204 for r in 0..n {
1205 for h in 0..nh {
1206 let dst = r * qrows + 2 * h * hd;
1207 let src = r * qdim + h * hd;
1208 dq[dst..dst + hd].copy_from_slice(&dqpre[src..src + hd]);
1209 dq[dst + hd..dst + 2 * hd].copy_from_slice(&dgate[src..src + hd]);
1210 }
1211 }
1212 dq
1213 } else {
1214 dqpre
1215 };
1216
1217 let mut dn1 = vec![0f32; n * hsz];
1219 ops::gemm_dx(&dqraw, wq, &mut dn1, n, hsz, qrows, pool);
1220 ops::gemm_dx(&dkpre, wk, &mut dn1, n, hsz, kvdim, pool);
1221 ops::gemm_dx(&dvproj, wv, &mut dn1, n, hsz, kvdim, pool);
1222 dn1
1223 }
1224
1225 fn gdn_attn_bwd(
1229 &self,
1230 attn: &FcdAttn,
1231 acts: &AttnActs,
1232 dattn: &[f32],
1233 b: usize,
1234 t: usize,
1235 ) -> Vec<f32> {
1236 let FcdAttn::Gdn { wqkv, wz, wa, wb, conv, a_log, dt_bias, norm, wout } = attn else {
1237 unreachable!("gdn_attn_bwd on a non-GDN layer");
1238 };
1239 let AttnActs::Gdn { qkv, z, a, b: bstr } = acts else {
1240 unreachable!("acts mismatch");
1241 };
1242 let d = self.gdn.expect("gdn layer without gdn dims");
1243 let (hsz, n) = (self.hidden, b * t);
1244 let pool = self.pool.as_deref();
1245 let (c_dim, vd, nv) = (d.c_dim(), d.vd(), d.nv);
1246
1247 let mut dof = vec![0f32; n * vd];
1248 ops::gemm_dx(dattn, wout, &mut dof, n, vd, hsz, pool);
1249
1250 let cfg = ops::GdnSeqCfg {
1251 nv: d.nv,
1252 nk: d.nk,
1253 dk: d.dk,
1254 dv: d.dv,
1255 kk: d.kk,
1256 rms_eps: self.eps,
1257 conv,
1258 a_log,
1259 dt_bias,
1260 norm,
1261 };
1262 let qkv64: Vec<f64> = qkv.iter().map(|&v| v as f64).collect();
1263 let z64: Vec<f64> = z.iter().map(|&v| v as f64).collect();
1264 let a64: Vec<f64> = a.iter().map(|&v| v as f64).collect();
1265 let b64: Vec<f64> = bstr.iter().map(|&v| v as f64).collect();
1266 let dof64: Vec<f64> = dof.iter().map(|&v| v as f64).collect();
1267 let mut pre64 = vec![0f64; n * c_dim];
1268 let mut cq64 = vec![0f64; n * c_dim];
1269 for bi in 0..b {
1270 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1271 ops::gdn_conv_fwd(
1272 &qkv64[r.clone()],
1273 t,
1274 c_dim,
1275 d.kk,
1276 conv,
1277 &mut pre64[r.clone()],
1278 &mut cq64[r],
1279 );
1280 }
1281
1282 let mut dcq64 = vec![0f64; n * c_dim];
1283 let mut dz64 = vec![0f64; n * vd];
1284 let mut da64 = vec![0f64; n * nv];
1285 let mut db64 = vec![0f64; n * nv];
1286 {
1287 let units = b * d.nk;
1288 let rep_v = d.nv / d.nk;
1289 let kd = d.nk * d.dk;
1290 let dcqp = SendMut(dcq64.as_mut_ptr());
1291 let dzp = SendMut(dz64.as_mut_ptr());
1292 let dap = SendMut(da64.as_mut_ptr());
1293 let dbp = SendMut(db64.as_mut_ptr());
1294 let (cqr, zr, ar, br, dor) = (&cq64, &z64, &a64, &b64, &dof64);
1295 let cfg_ref = &cfg;
1296 let run_unit = |u: usize| {
1297 let (bi, ko) = (u / d.nk, u % d.nk);
1298 let mut dcq_l = vec![0f64; t * c_dim];
1301 let mut dz_l = vec![0f64; t * vd];
1302 let mut da_l = vec![0f64; t * nv];
1303 let mut db_l = vec![0f64; t * nv];
1304 ops::gdn_group_bwd(
1305 &cqr[bi * t * c_dim..(bi + 1) * t * c_dim],
1306 &zr[bi * t * vd..(bi + 1) * t * vd],
1307 &ar[bi * t * nv..(bi + 1) * t * nv],
1308 &br[bi * t * nv..(bi + 1) * t * nv],
1309 t,
1310 cfg_ref,
1311 ko,
1312 &dor[bi * t * vd..(bi + 1) * t * vd],
1313 &mut dcq_l,
1314 &mut dz_l,
1315 &mut da_l,
1316 &mut db_l,
1317 );
1318 for p in 0..t {
1321 let row = (bi * t + p) * c_dim;
1322 for c in ko * d.dk..(ko + 1) * d.dk {
1323 unsafe {
1324 *dcqp.at(row + c) = dcq_l[p * c_dim + c];
1325 *dcqp.at(row + kd + c) = dcq_l[p * c_dim + kd + c];
1326 }
1327 }
1328 for hh in 0..rep_v {
1329 let h = ko * rep_v + hh;
1330 for dj in 0..d.dv {
1331 unsafe {
1332 *dcqp.at(row + 2 * kd + h * d.dv + dj) =
1333 dcq_l[p * c_dim + 2 * kd + h * d.dv + dj];
1334 *dzp.at((bi * t + p) * vd + h * d.dv + dj) =
1335 dz_l[p * vd + h * d.dv + dj];
1336 }
1337 }
1338 unsafe {
1339 *dap.at((bi * t + p) * nv + h) = da_l[p * nv + h];
1340 *dbp.at((bi * t + p) * nv + h) = db_l[p * nv + h];
1341 }
1342 }
1343 }
1344 };
1345 match pool {
1346 Some(p) if units > 1 => p.run(&|widx, nw| {
1347 for u in (widx..units).step_by(nw) {
1348 run_unit(u);
1349 }
1350 }),
1351 _ => {
1352 for u in 0..units {
1353 run_unit(u);
1354 }
1355 }
1356 }
1357 }
1358
1359 let mut dqkv64 = vec![0f64; n * c_dim];
1360 for bi in 0..b {
1361 let r = bi * t * c_dim..(bi + 1) * t * c_dim;
1362 ops::gdn_conv_bwd(
1363 &pre64[r.clone()],
1364 t,
1365 c_dim,
1366 d.kk,
1367 conv,
1368 &dcq64[r.clone()],
1369 &mut dqkv64[r],
1370 );
1371 }
1372 let to32 = |v: &[f64]| -> Vec<f32> { v.iter().map(|&x| x as f32).collect() };
1373 let (dqkv, dz, da, db) = (to32(&dqkv64), to32(&dz64), to32(&da64), to32(&db64));
1374
1375 let mut dn1 = vec![0f32; n * hsz];
1376 ops::gemm_dx(&dqkv, wqkv, &mut dn1, n, hsz, c_dim, pool);
1377 ops::gemm_dx(&dz, wz, &mut dn1, n, hsz, vd, pool);
1378 ops::gemm_dx(&da, wa, &mut dn1, n, hsz, nv, pool);
1379 ops::gemm_dx(&db, wb, &mut dn1, n, hsz, nv, pool);
1380 dn1
1381 }
1382
1383 fn forward_hidden(
1388 &self,
1389 ids: &[u32],
1390 b: usize,
1391 t: usize,
1392 ts: Option<&TrainState>,
1393 student: bool,
1394 mut keep: Option<&mut Vec<Vec<f32>>>,
1395 ) -> Vec<f32> {
1396 let hsz = self.hidden;
1397 let mut h = vec![0f32; b * t * hsz];
1398 for (r, &id) in ids.iter().enumerate() {
1399 let src = (id as usize).min(self.embed.len() / hsz - 1) * hsz;
1400 h[r * hsz..(r + 1) * hsz].copy_from_slice(&self.embed[src..src + hsz]);
1401 }
1402 for li in 0..self.nl {
1403 if let Some(k) = keep.as_deref_mut() {
1404 k.push(h.clone());
1405 }
1406 let wts = ln_ffn(self, if student { ts } else { None }, li);
1407 let nys = student && self.o1_flags[li];
1408 h = self.layer_forward(li, &h, b, t, &wts, nys, false).0;
1409 }
1410 h
1411 }
1412
1413 fn loss_and_dhidden(
1416 &self,
1417 hs: &[f32],
1418 ht: &[f32],
1419 targets: &[u32],
1420 kl_w: f64,
1421 ) -> (f64, f64, Vec<f32>) {
1422 let hsz = self.hidden;
1423 let n = targets.len();
1424 let pool = self.pool.as_deref();
1425 let wh = self.head_weight();
1426 let vs = self.vocab;
1427
1428 let mut ns = vec![0f32; n * hsz];
1429 let mut invs = vec![0f32; n];
1430 ops::rmsnorm_fwd(hs, &self.final_norm, self.eps, self.gemma, &mut ns, &mut invs);
1431 let mut nt = vec![0f32; n * hsz];
1432 let mut invt = vec![0f32; n];
1433 ops::rmsnorm_fwd(ht, &self.final_norm, self.eps, self.gemma, &mut nt, &mut invt);
1434
1435 let inv_n = 1.0 / n as f64;
1436 let mut ce_sum = 0f64;
1437 let mut kl_sum = 0f64;
1438 let mut dns = vec![0f32; n * hsz];
1439 let mut ls = vec![0f32; LM_CHUNK * vs];
1440 let mut lt = vec![0f32; LM_CHUNK * vs];
1441 let mut dlg = vec![0f32; LM_CHUNK * vs];
1442 let mut r0 = 0usize;
1443 while r0 < n {
1444 let r1 = (r0 + LM_CHUNK).min(n);
1445 let c = r1 - r0;
1446 ops::gemm_nt(&ns[r0 * hsz..r1 * hsz], wh, &mut ls[..c * vs], c, hsz, vs, pool);
1447 ops::gemm_nt(&nt[r0 * hsz..r1 * hsz], wh, &mut lt[..c * vs], c, hsz, vs, pool);
1448 for r in 0..c {
1449 let (ce, kl) = ops::ce_kl_position(
1450 &ls[r * vs..(r + 1) * vs],
1451 <[r * vs..(r + 1) * vs],
1452 targets[r0 + r] as usize,
1453 kl_w,
1454 inv_n,
1455 &mut dlg[r * vs..(r + 1) * vs],
1456 );
1457 ce_sum += ce;
1458 kl_sum += kl;
1459 }
1460 ops::gemm_dx(
1461 &dlg[..c * vs],
1462 wh,
1463 &mut dns[r0 * hsz..r1 * hsz],
1464 c,
1465 hsz,
1466 vs,
1467 pool,
1468 );
1469 r0 = r1;
1470 }
1471
1472 let mut dhs = vec![0f32; n * hsz];
1473 ops::rmsnorm_bwd(hs, &self.final_norm, &invs, &dns, self.gemma, &mut dhs, None);
1474 (ce_sum * inv_n, kl_sum * inv_n, dhs)
1475 }
1476
1477 fn backward(
1480 &self,
1481 b: usize,
1482 t: usize,
1483 keep: &[Vec<f32>],
1484 dh_last: Vec<f32>,
1485 ts: &mut TrainState,
1486 ) {
1487 let TrainState { layers, data, grad, .. } = ts;
1490 let mut dh = dh_last;
1491 for li in (0..self.nl).rev() {
1492 let h_in = &keep[li];
1493 let nys = self.o1_flags[li];
1494 let slot = layers.iter().position(|&x| x == li);
1495 let wts = match slot {
1496 Some(s) => {
1497 let bi = s * PARAMS_PER_LAYER;
1498 LnFfn {
1499 iln: &data[bi],
1500 pln: &data[bi + 1],
1501 gate: &data[bi + 2],
1502 up: &data[bi + 3],
1503 down: &data[bi + 4],
1504 }
1505 }
1506 None => {
1507 let l = &self.layers[li];
1508 LnFfn { iln: &l.iln, pln: &l.pln, gate: &l.gate, up: &l.up, down: &l.down }
1509 }
1510 };
1511 let (_, acts) = self.layer_forward(li, h_in, b, t, &wts, nys, true);
1512 let acts = acts.expect("want_acts");
1513 dh = match slot {
1514 Some(s) => {
1515 let gb = s * PARAMS_PER_LAYER;
1516 let gr = &mut grad[gb..gb + PARAMS_PER_LAYER];
1517 self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, Some(gr))
1518 }
1519 None => self.layer_backward(li, h_in, b, t, &wts, nys, &acts, &dh, None),
1520 };
1521 }
1522 }
1523
1524 #[doc(hidden)]
1531 pub fn loss_and_grads_for_test(
1532 &self,
1533 ids: &[u32],
1534 tgt: &[u32],
1535 b: usize,
1536 t: usize,
1537 ts: &mut TrainState,
1538 kl_w: f64,
1539 ) -> f64 {
1540 let ht = self.forward_hidden(ids, b, t, None, false, None);
1541 let mut keep = Vec::with_capacity(self.nl);
1542 let hs = self.forward_hidden(ids, b, t, Some(ts), true, Some(&mut keep));
1543 let (ce, kl, dhs) = self.loss_and_dhidden(&hs, &ht, tgt, kl_w);
1544 ts.zero_grad();
1545 self.backward(b, t, &keep, dhs, ts);
1546 (1.0 - kl_w) * ce + kl_w * kl
1547 }
1548
1549 pub fn val_ppl(
1553 &self,
1554 va: &[u32],
1555 ts: Option<&TrainState>,
1556 student: bool,
1557 bs: usize,
1558 nrounds: usize,
1559 seq: usize,
1560 ) -> f64 {
1561 let nwin = nrounds * bs;
1562 if va.len() < seq + 2 || nwin == 0 {
1563 return f64::NAN;
1564 }
1565 let stride = (va.len() - seq - 1) / nwin;
1566 let hsz = self.hidden;
1567 let wh = self.head_weight();
1568 let vs = self.vocab;
1569 let pool = self.pool.as_deref();
1570 let mut nll = 0f64;
1571 let mut cnt = 0usize;
1572 for j in 0..nrounds {
1573 let mut ids = Vec::with_capacity(bs * seq);
1574 let mut tgt = Vec::with_capacity(bs * seq);
1575 for bi in 0..bs {
1576 let off = ((j * bs + bi) * stride.max(1)).min(va.len() - seq - 1);
1577 ids.extend_from_slice(&va[off..off + seq]);
1578 tgt.extend_from_slice(&va[off + 1..off + seq + 1]);
1579 }
1580 let h = self.forward_hidden(&ids, bs, seq, ts, student, None);
1581 let n = bs * seq;
1582 let mut ns = vec![0f32; n * hsz];
1583 let mut inv = vec![0f32; n];
1584 ops::rmsnorm_fwd(&h, &self.final_norm, self.eps, self.gemma, &mut ns, &mut inv);
1585 let mut lg = vec![0f32; LM_CHUNK * vs];
1586 let mut r0 = 0usize;
1587 while r0 < n {
1588 let r1 = (r0 + LM_CHUNK).min(n);
1589 let c = r1 - r0;
1590 ops::gemm_nt(&ns[r0 * hsz..r1 * hsz], wh, &mut lg[..c * vs], c, hsz, vs, pool);
1591 for r in 0..c {
1592 let row = &lg[r * vs..(r + 1) * vs];
1593 let target = tgt[r0 + r] as usize;
1594 let mut mx = f64::NEG_INFINITY;
1595 for &v in row {
1596 mx = mx.max(v as f64);
1597 }
1598 let mut s = 0f64;
1599 for &v in row {
1600 s += (v as f64 - mx).exp();
1601 }
1602 nll += mx + s.ln() - row[target.min(vs - 1)] as f64;
1603 cnt += 1;
1604 }
1605 r0 = r1;
1606 }
1607 }
1608 (nll / cnt.max(1) as f64).exp()
1609 }
1610}
1611
1612pub fn run_polish(
1624 model: &Arc<CmfModel>,
1625 o1: &O1Cfg,
1626 hp: &FcdHyper,
1627 tr: &[u32],
1628 va: &[u32],
1629 out: &std::path::Path,
1630 gate: Option<&GenGateCfg>,
1631) -> Result<FcdReport, String> {
1632 if tr.len() < hp.seq + 2 {
1633 return Err(format!(
1634 "train corpus too small: {} tokens < seq+2 = {}",
1635 tr.len(),
1636 hp.seq + 2
1637 ));
1638 }
1639 let fm = FcdModel::from_cmf(model, o1)?;
1640 let converted = fm.converted();
1641 if converted.is_empty() {
1642 return Err("no converted layers under this --o1 spec (nothing to polish)".into());
1643 }
1644 tracing::info!(
1645 "fcd: {} layers converted ({} trainable tensors), m={} w={} sink={}, \
1646 corpus train {} / val {} tokens",
1647 converted.len(),
1648 converted.len() * PARAMS_PER_LAYER,
1649 fm.nys.m,
1650 fm.nys.w,
1651 fm.nys.sink,
1652 tr.len(),
1653 va.len()
1654 );
1655
1656 let mut ts = TrainState::new(&fm);
1657 let teacher_ppl = fm.val_ppl(va, None, false, hp.bs, 2, hp.seq);
1658 let ppl_start = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
1659 tracing::info!(
1660 "fcd: quick-val teacher ppl {teacher_ppl:.2} | zero-shot o1 student ppl {ppl_start:.2}"
1661 );
1662
1663 let mut gate_state: Option<(Pipeline, Vec<f64>)> = match gate {
1665 Some(g) if !g.prompts.is_empty() => {
1666 let greedy = SamplerConfig {
1667 temperature: 0.0,
1668 top_p: 1.0,
1669 top_k: 0,
1670 repetition_penalty: 1.0,
1671 min_p: 0.0,
1672 seed: Some(0),
1673 };
1674 let mut pipe = Pipeline::from_model(model, greedy)
1675 .map_err(|e| format!("gen-gate pipeline: {e}"))?;
1676 pipe.set_o1(Some(o1.clone()));
1677 apply_trainables(&mut pipe, &fm, &ts);
1678 let base = gate_gen_scores(&mut pipe, g)?;
1679 tracing::info!("fcd gen-gate baseline loop-scores: {base:?}");
1680 Some((pipe, base))
1681 }
1682 Some(_) => {
1683 tracing::warn!("fcd gen-gate requested but val stream too short — gate off");
1684 None
1685 }
1686 None => None,
1687 };
1688 let init_snapshot: Option<Vec<Vec<f32>>> =
1690 gate_state.is_some().then(|| ts.data.clone());
1691 let mut gate_evals: Vec<(usize, f64, Vec<f64>, bool)> = Vec::new();
1692
1693 let mut rng = SplitMix64::new(hp.seed);
1694 let mut best: (f64, Option<Vec<Vec<f32>>>, usize) = (f64::INFINITY, None, 0);
1695 let mut losses: Vec<(f64, f64)> = Vec::with_capacity(hp.steps);
1696 let t0 = std::time::Instant::now();
1697 let n_per_step = hp.bs * hp.seq;
1698 for st in 1..=hp.steps {
1699 let mut ids = Vec::with_capacity(n_per_step);
1702 let mut tgt = Vec::with_capacity(n_per_step);
1703 for _ in 0..hp.bs {
1704 let off = (rng.next_u64() as usize) % (tr.len() - hp.seq - 1);
1705 ids.extend_from_slice(&tr[off..off + hp.seq]);
1706 tgt.extend_from_slice(&tr[off + 1..off + hp.seq + 1]);
1707 }
1708
1709 let ht = fm.forward_hidden(&ids, hp.bs, hp.seq, None, false, None);
1710 let mut keep: Vec<Vec<f32>> = Vec::with_capacity(fm.nl);
1711 let hs = fm.forward_hidden(&ids, hp.bs, hp.seq, Some(&ts), true, Some(&mut keep));
1712 let (ce, kl, dhs) = fm.loss_and_dhidden(&hs, &ht, &tgt, hp.kl_w);
1713 ts.zero_grad();
1714 fm.backward(hp.bs, hp.seq, &keep, dhs, &mut ts);
1715 let gn = ts.clip_and_step(hp.lr);
1716 losses.push((ce, kl));
1717
1718 let el = t0.elapsed().as_secs_f64();
1719 tracing::info!(
1720 "fcd step {st}/{}: ce {ce:.3} kl {kl:.3} |g| {gn:.3} ({:.1}s/step)",
1721 hp.steps,
1722 el / st as f64
1723 );
1724 if hp.eval_every > 0 && st % hp.eval_every == 0 {
1725 let p = fm.val_ppl(va, Some(&ts), true, hp.bs, 2, hp.seq);
1726 match (&mut gate_state, gate) {
1727 (Some((pipe, base)), Some(g)) => {
1728 apply_trainables(pipe, &fm, &ts);
1729 let scores = gate_gen_scores(pipe, g)?;
1730 let pass = gate_pass(&scores, base, g.threshold, g.baseline_slack);
1731 let tag = if pass && p < best.0 {
1732 best = (p, Some(ts.data.clone()), st);
1733 " *best*"
1734 } else {
1735 ""
1736 };
1737 tracing::info!(
1738 "fcd eval step {st}: val ppl {p:.2} | gen-gate {} (loop-scores {scores:?}){tag}",
1739 if pass { "PASS" } else { "FAIL" }
1740 );
1741 gate_evals.push((st, p, scores, pass));
1742 }
1743 _ => {
1744 let tag = if p < best.0 {
1745 best = (p, Some(ts.data.clone()), st);
1746 " *best*"
1747 } else {
1748 ""
1749 };
1750 tracing::info!("fcd eval step {st}: val ppl {p:.2}{tag}");
1751 }
1752 }
1753 }
1754 }
1755
1756 let mut gate_chosen: Option<usize> = None;
1760 if let Some(snap) = best.1.take() {
1761 ts.data = snap;
1762 gate_chosen = Some(best.2);
1763 tracing::info!(
1764 "fcd: restored best checkpoint from step {} (val ppl {:.2})",
1765 best.2,
1766 best.0
1767 );
1768 } else if let Some(init) = init_snapshot {
1769 ts.data = init;
1770 tracing::info!(
1771 "fcd: polish rejected by generation gate — identity artifact (zero-shot state written; claim 13 floor)"
1772 );
1773 }
1774 let ppl_final = fm.val_ppl(va, Some(&ts), true, hp.bs, 6, hp.seq);
1775 let report = FcdReport {
1776 converted: converted.clone(),
1777 teacher_ppl,
1778 ppl_start,
1779 ppl_best: best.0.min(ppl_final),
1780 best_step: best.2,
1781 ppl_final,
1782 steps_run: hp.steps,
1783 sec_per_step: t0.elapsed().as_secs_f64() / hp.steps.max(1) as f64,
1784 losses,
1785 gate: gate_state.map(|(_, base)| GateReport {
1786 baseline: base,
1787 evals: gate_evals,
1788 chosen: gate_chosen,
1789 }),
1790 };
1791 save_polished(model, out, &fm, &ts, o1, hp, &report)?;
1792 Ok(report)
1793}
1794
1795fn apply_trainables(pipe: &mut Pipeline, fm: &FcdModel, ts: &TrainState) {
1799 let hidden = fm.hidden;
1800 for (slot, &li) in ts.layers.iter().enumerate() {
1801 let b = slot * PARAMS_PER_LAYER;
1802 let inter = fm.layers[li].inter;
1803 let lw = &mut pipe.weights.layers[li];
1804 lw.input_norm = ts.data[b].clone();
1805 lw.post_norm = ts.data[b + 1].clone();
1806 lw.ffn = FfnKind::Dense(DenseFfn {
1807 gate_proj: QTensor::from_f32(ts.data[b + 2].clone(), inter, hidden),
1808 up_proj: QTensor::from_f32(ts.data[b + 3].clone(), inter, hidden),
1809 down_proj: QTensor::from_f32(ts.data[b + 4].clone(), hidden, inter),
1810 });
1811 }
1812}
1813
1814fn gate_gen_scores(pipe: &mut Pipeline, g: &GenGateCfg) -> Result<Vec<f64>, String> {
1816 g.prompts
1817 .iter()
1818 .map(|p| {
1819 pipe.generate_from_ids(p, g.gen_tokens, None, None)
1820 .map(|r| loop_score(&r.token_ids))
1821 })
1822 .collect()
1823}
1824
1825fn save_polished(
1830 model: &CmfModel,
1831 out: &std::path::Path,
1832 fm: &FcdModel,
1833 ts: &TrainState,
1834 o1: &O1Cfg,
1835 hp: &FcdHyper,
1836 report: &FcdReport,
1837) -> Result<(), String> {
1838 use cortiq_core::format::TensorSpec;
1839 let mut replace: std::collections::HashMap<String, (usize, usize)> =
1840 std::collections::HashMap::new(); for (s, &li) in ts.layers.iter().enumerate() {
1842 let p = format!("model.layers.{li}.");
1843 for (k, suffix) in [
1844 (0usize, "input_layernorm.weight"),
1845 (1, "post_attention_layernorm.weight"),
1846 (2, "mlp.gate_proj.weight"),
1847 (3, "mlp.up_proj.weight"),
1848 (4, "mlp.down_proj.weight"),
1849 ] {
1850 replace.insert(format!("{p}{suffix}"), (s, k));
1851 }
1852 }
1853 let mut specs = Vec::with_capacity(model.tensors.len());
1854 for t in &model.tensors {
1855 if let Some(&(s, k)) = replace.get(&t.name) {
1856 let data = &ts.data[s * PARAMS_PER_LAYER + k];
1857 let mut bytes = Vec::with_capacity(data.len() * 4);
1858 for v in data {
1859 bytes.extend_from_slice(&v.to_le_bytes());
1860 }
1861 specs.push(TensorSpec {
1862 name: t.name.clone(),
1863 dtype: TensorDtype::F32,
1864 shape: t.shape.clone(),
1865 data: bytes,
1866 });
1867 } else {
1868 specs.push(TensorSpec {
1869 name: t.name.clone(),
1870 dtype: t.dtype,
1871 shape: t.shape.clone(),
1872 data: model.entry_bytes(t).to_vec(),
1873 });
1874 }
1875 }
1876
1877 let mut header = model.header.clone();
1878 let mut prov = match header.provenance.take() {
1879 Some(serde_json::Value::Object(m)) => m,
1880 _ => serde_json::Map::new(),
1881 };
1882 let layers_json = match &o1.layers {
1883 O1Layers::All => serde_json::json!("all"),
1884 O1Layers::Deep(n) => serde_json::json!(format!("deep{n}")),
1885 O1Layers::List(v) => serde_json::json!(v),
1886 };
1887 prov.insert(
1888 "o1_attn".into(),
1889 serde_json::json!({
1890 "layers": layers_json, "m": o1.m, "w": o1.w, "sink": o1.sink
1891 }),
1892 );
1893 prov.insert(
1894 "fcd".into(),
1895 serde_json::json!({
1896 "steps": hp.steps, "lr": hp.lr, "kl_w": hp.kl_w,
1897 "bs": hp.bs, "seq": hp.seq,
1898 "teacher_ppl": report.teacher_ppl,
1899 "ppl_start": report.ppl_start,
1900 "ppl_final": report.ppl_final,
1901 "best_step": report.best_step,
1902 "converted_layers": report.converted,
1903 }),
1904 );
1905 header.provenance = Some(serde_json::Value::Object(prov));
1906 let _ = fm; let masks = if model.masks.masks.is_empty() {
1909 None
1910 } else {
1911 Some(&model.masks)
1912 };
1913 CmfModel::write(out, &header, &specs, masks, model.vocab.as_deref())
1914 .map_err(|e| format!("writing polished cmf: {e}"))
1915}
1916
1917#[cfg(test)]
1918mod tests {
1919 use super::*;
1920
1921 #[test]
1923 fn gate_selects_lowest_ppl_among_passing() {
1924 let base = vec![0.10, 0.00, 0.20];
1925 let evals = vec![
1926 (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]), ];
1931 let sel = select_checkpoint(&evals, &base, 0.35, 0.10);
1932 assert_eq!(sel, Some(2), "step 75 is the lowest-ppl PASSING checkpoint");
1933 }
1934
1935 #[test]
1938 fn gate_all_fail_is_identity() {
1939 let base = vec![0.0, 0.0, 0.0];
1940 let evals = vec![
1941 (25usize, 15.0, vec![0.50, 0.0, 0.0]),
1942 (50, 14.0, vec![0.0, 0.36, 0.0]),
1943 (75, 13.0, vec![0.0, 0.0, 0.11]), ];
1945 assert_eq!(select_checkpoint(&evals, &base, 0.35, 0.10), None);
1946 }
1947
1948 #[test]
1951 fn gate_boundaries_and_tie_break() {
1952 let base = vec![0.25];
1953 assert!(gate_pass(&[0.35], &base, 0.35, 0.10), "== threshold passes");
1954 assert!(gate_pass(&[0.35], &[0.25], 0.35, 0.10), "== base+slack passes");
1955 assert!(!gate_pass(&[0.351], &base, 0.35, 0.10));
1956 assert!(!gate_pass(&[0.30], &[0.10], 0.35, 0.10), "0.30 > 0.10+0.10");
1957 let evals = vec![
1958 (25usize, 20.0, vec![0.10]),
1959 (50, 20.0, vec![0.10]),
1960 ];
1961 assert_eq!(
1962 select_checkpoint(&evals, &base, 0.35, 0.10),
1963 Some(0),
1964 "equal ppl → earliest checkpoint"
1965 );
1966 }
1967}