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