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