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