1use crate::backends::llm_policy::OptimizerKind;
7use anyhow::{Context, Result, bail};
8use serde_json::json;
9use std::fs::File;
10use std::io::Write;
11use std::path::Path;
12use std::time::Instant;
13use wide::f32x8;
14
15use super::kernel;
16use super::profiling::{NullProfiler, ProfilerSink};
17use super::tensor::Tensor1D;
18use super::weights::Weights;
19
20#[derive(Debug, Clone)]
22pub struct Config {
23 pub vocab_size: usize,
25 pub hidden_size: usize,
27 pub num_layers: usize,
29 pub num_heads: usize,
31 pub head_dim: usize,
33 pub intermediate_size: usize,
35 pub layer_norm_eps: f32,
37 pub group_norm_eps: f32,
39
40 pub decay_low_rank: usize, pub a_low_rank: usize,
44 pub v_low_rank: usize,
46 pub g_low_rank: usize,
48}
49
50impl Default for Config {
51 fn default() -> Self {
52 Self {
53 vocab_size: 256,
54 hidden_size: 256,
55 num_layers: 12,
56 num_heads: 4, head_dim: 64,
58 intermediate_size: 1024,
59 layer_norm_eps: 1e-5,
60 group_norm_eps: 64e-5,
61 decay_low_rank: 32,
62 a_low_rank: 32,
63 v_low_rank: 32,
64 g_low_rank: 64,
65 }
66 }
67}
68
69impl Config {
70 pub fn validate(&self) -> Result<()> {
72 if self.vocab_size == 0 {
73 bail!("rwkv7 vocab_size must be > 0");
74 }
75 if self.head_dim != 64 {
76 bail!("rwkv7 head_dim must be 64 for current kernels");
77 }
78 if self.hidden_size != self.num_heads * self.head_dim {
79 bail!(
80 "rwkv7 hidden_size must equal num_heads * head_dim ({} != {} * {})",
81 self.hidden_size,
82 self.num_heads,
83 self.head_dim
84 );
85 }
86 if self.num_layers == 0 {
87 bail!("rwkv7 num_layers must be > 0");
88 }
89 if self.intermediate_size == 0 {
90 bail!("rwkv7 intermediate_size must be > 0");
91 }
92 Ok(())
93 }
94}
95
96#[derive(Clone)]
98pub struct LayerState {
99 pub att_x_prev: Tensor1D,
101 pub att_state: Tensor1D, pub ffn_x_prev: Tensor1D,
105}
106
107impl LayerState {
108 fn new(cfg: &Config) -> Self {
109 let state_size = cfg.num_heads * cfg.head_dim * cfg.head_dim;
110 Self {
111 att_x_prev: Tensor1D::zeros(cfg.hidden_size),
112 att_state: Tensor1D::zeros(state_size),
113 ffn_x_prev: Tensor1D::zeros(cfg.hidden_size),
114 }
115 }
116
117 fn copy_from(&mut self, other: &Self) {
118 self.att_x_prev.copy_from(&other.att_x_prev);
119 self.att_state.copy_from(&other.att_state);
120 self.ffn_x_prev.copy_from(&other.ffn_x_prev);
121 }
122}
123
124#[derive(Clone)]
126pub struct State {
127 pub layers: Vec<LayerState>,
129 pub v_first: Tensor1D,
131 pub v_first_set: bool,
133}
134
135impl State {
136 pub fn new(cfg: &Config) -> Self {
138 Self {
139 layers: (0..cfg.num_layers).map(|_| LayerState::new(cfg)).collect(),
140 v_first: Tensor1D::zeros(cfg.hidden_size),
141 v_first_set: false,
142 }
143 }
144
145 pub fn reset(&mut self) {
147 self.v_first_set = false;
148 self.v_first.zero();
149 for layer in &mut self.layers {
150 layer.att_x_prev.zero();
151 layer.att_state.zero();
152 layer.ffn_x_prev.zero();
153 }
154 }
155
156 pub(crate) fn copy_from(&mut self, other: &Self) {
157 debug_assert_eq!(self.layers.len(), other.layers.len());
158 self.v_first.clone_from(&other.v_first);
159 self.v_first_set = other.v_first_set;
160 for (dst, src) in self.layers.iter_mut().zip(other.layers.iter()) {
161 dst.copy_from(src);
162 }
163 }
164}
165
166#[derive(Clone)]
168struct AttentionWeights {
169 x_r: Tensor1D,
171 x_w: Tensor1D,
172 x_k: Tensor1D,
173 x_v: Tensor1D,
174 x_a: Tensor1D,
175 x_g: Tensor1D,
176
177 rkv_proj: Tensor1D,
180
181 o_proj: Tensor1D,
183
184 w1: Tensor1D, w2: Tensor1D, w0: Tensor1D, a1: Tensor1D, a2: Tensor1D, a0: Tensor1D, v1: Option<Tensor1D>, v2: Option<Tensor1D>, v0: Option<Tensor1D>, g1: Tensor1D, g2: Tensor1D, k_k: Tensor1D, k_a: Tensor1D, r_k: Tensor1D, g_norm_w: Tensor1D, g_norm_b: Tensor1D, }
212
213#[derive(Clone)]
215struct FfnWeights {
216 x_k: Tensor1D, key_w: Tensor1D, value_w: Tensor1D, }
220
221#[derive(Clone)]
223struct BlockWeights {
224 pre_norm_w: Option<Tensor1D>,
226 pre_norm_b: Option<Tensor1D>,
227
228 attn_norm_w: Tensor1D,
230 attn_norm_b: Tensor1D,
231
232 ffn_norm_w: Tensor1D,
234 ffn_norm_b: Tensor1D,
235
236 attn: AttentionWeights,
237 ffn: FfnWeights,
238}
239
240#[derive(Clone)]
242pub struct Model {
243 cfg: Config,
244
245 embeddings: Tensor1D,
247
248 ln_out_w: Tensor1D,
250 ln_out_b: Tensor1D,
251
252 lm_head: Tensor1D,
254
255 blocks: Vec<BlockWeights>,
257}
258
259#[derive(Clone)]
260struct AdamTensorState {
261 m: Tensor1D,
262 v: Tensor1D,
263}
264
265impl AdamTensorState {
266 #[inline]
267 fn new(len: usize) -> Self {
268 Self {
269 m: Tensor1D::zeros(len),
270 v: Tensor1D::zeros(len),
271 }
272 }
273}
274
275#[derive(Clone)]
276struct AttentionAdamState {
277 x_r: AdamTensorState,
278 x_w: AdamTensorState,
279 x_k: AdamTensorState,
280 x_v: AdamTensorState,
281 x_a: AdamTensorState,
282 x_g: AdamTensorState,
283 rkv_proj: AdamTensorState,
284 o_proj: AdamTensorState,
285 w1: AdamTensorState,
286 w2: AdamTensorState,
287 w0: AdamTensorState,
288 a1: AdamTensorState,
289 a2: AdamTensorState,
290 a0: AdamTensorState,
291 v1: Option<AdamTensorState>,
292 v2: Option<AdamTensorState>,
293 v0: Option<AdamTensorState>,
294 g1: AdamTensorState,
295 g2: AdamTensorState,
296 k_k: AdamTensorState,
297 k_a: AdamTensorState,
298 r_k: AdamTensorState,
299 g_norm_w: AdamTensorState,
300 g_norm_b: AdamTensorState,
301}
302
303#[derive(Clone)]
304struct FfnAdamState {
305 x_k: AdamTensorState,
306 key_w: AdamTensorState,
307 value_w: AdamTensorState,
308}
309
310#[derive(Clone)]
311struct BlockAdamState {
312 pre_norm_w: Option<AdamTensorState>,
313 pre_norm_b: Option<AdamTensorState>,
314 attn_norm_w: AdamTensorState,
315 attn_norm_b: AdamTensorState,
316 ffn_norm_w: AdamTensorState,
317 ffn_norm_b: AdamTensorState,
318 attn: AttentionAdamState,
319 ffn: FfnAdamState,
320}
321
322#[derive(Clone)]
323pub struct FullAdamState {
325 embeddings: AdamTensorState,
326 ln_out_w: AdamTensorState,
327 ln_out_b: AdamTensorState,
328 lm_head: AdamTensorState,
329 blocks: Vec<BlockAdamState>,
330}
331
332#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
333pub struct TrainScopeMask {
335 pub embed: bool,
337 pub pre_norm: bool,
339 pub attn_norm: bool,
341 pub ffn_norm: bool,
343 pub attn: bool,
345 pub ffn: bool,
347 pub head: bool,
349 pub bias: bool,
351}
352
353impl TrainScopeMask {
354 #[inline]
355 pub fn all() -> Self {
357 Self {
358 embed: true,
359 pre_norm: true,
360 attn_norm: true,
361 ffn_norm: true,
362 attn: true,
363 ffn: true,
364 head: true,
365 bias: true,
366 }
367 }
368
369 #[inline]
370 pub fn trains_non_head_params(&self) -> bool {
372 self.embed || self.pre_norm || self.attn_norm || self.ffn_norm || self.attn || self.ffn
373 }
374
375 #[inline]
376 pub fn trains_any_params(&self) -> bool {
378 self.trains_non_head_params() || self.head || self.bias
379 }
380}
381
382#[derive(Clone)]
383struct AttentionGradState {
384 x_r: Tensor1D,
385 x_w: Tensor1D,
386 x_k: Tensor1D,
387 x_v: Tensor1D,
388 x_a: Tensor1D,
389 x_g: Tensor1D,
390 rkv_proj: Tensor1D,
391 o_proj: Tensor1D,
392 w1: Tensor1D,
393 w2: Tensor1D,
394 w0: Tensor1D,
395 a1: Tensor1D,
396 a2: Tensor1D,
397 a0: Tensor1D,
398 v1: Option<Tensor1D>,
399 v2: Option<Tensor1D>,
400 v0: Option<Tensor1D>,
401 g1: Tensor1D,
402 g2: Tensor1D,
403 k_k: Tensor1D,
404 k_a: Tensor1D,
405 r_k: Tensor1D,
406 g_norm_w: Tensor1D,
407 g_norm_b: Tensor1D,
408}
409
410#[derive(Clone)]
411struct FfnGradState {
412 x_k: Tensor1D,
413 key_w: Tensor1D,
414 value_w: Tensor1D,
415}
416
417#[derive(Clone)]
418struct BlockGradState {
419 pre_norm_w: Option<Tensor1D>,
420 pre_norm_b: Option<Tensor1D>,
421 attn_norm_w: Tensor1D,
422 attn_norm_b: Tensor1D,
423 ffn_norm_w: Tensor1D,
424 ffn_norm_b: Tensor1D,
425 attn: AttentionGradState,
426 ffn: FfnGradState,
427}
428
429#[derive(Clone)]
430struct FullGradState {
431 embeddings: Tensor1D,
432 ln_out_w: Tensor1D,
433 ln_out_b: Tensor1D,
434 lm_head: Tensor1D,
435 blocks: Vec<BlockGradState>,
436}
437
438impl FullGradState {
439 fn zero(&mut self) {
440 self.embeddings.zero();
441 self.ln_out_w.zero();
442 self.ln_out_b.zero();
443 self.lm_head.zero();
444 for block in &mut self.blocks {
445 if let Some(t) = block.pre_norm_w.as_mut() {
446 t.zero();
447 }
448 if let Some(t) = block.pre_norm_b.as_mut() {
449 t.zero();
450 }
451 block.attn_norm_w.zero();
452 block.attn_norm_b.zero();
453 block.ffn_norm_w.zero();
454 block.ffn_norm_b.zero();
455
456 block.attn.x_r.zero();
457 block.attn.x_w.zero();
458 block.attn.x_k.zero();
459 block.attn.x_v.zero();
460 block.attn.x_a.zero();
461 block.attn.x_g.zero();
462 block.attn.rkv_proj.zero();
463 block.attn.o_proj.zero();
464 block.attn.w1.zero();
465 block.attn.w2.zero();
466 block.attn.w0.zero();
467 block.attn.a1.zero();
468 block.attn.a2.zero();
469 block.attn.a0.zero();
470 if let Some(t) = block.attn.v1.as_mut() {
471 t.zero();
472 }
473 if let Some(t) = block.attn.v2.as_mut() {
474 t.zero();
475 }
476 if let Some(t) = block.attn.v0.as_mut() {
477 t.zero();
478 }
479 block.attn.g1.zero();
480 block.attn.g2.zero();
481 block.attn.k_k.zero();
482 block.attn.k_a.zero();
483 block.attn.r_k.zero();
484 block.attn.g_norm_w.zero();
485 block.attn.g_norm_b.zero();
486
487 block.ffn.x_k.zero();
488 block.ffn.key_w.zero();
489 block.ffn.value_w.zero();
490 }
491 }
492}
493
494struct AdamStep {
495 lr: f32,
496 clip: f32,
497 b1: f32,
498 b2: f32,
499 eps: f32,
500 bias_corr1: f32,
501 bias_corr2: f32,
502}
503
504#[derive(Clone)]
505struct LayerTrainTrace {
506 x_in: Tensor1D,
507 x_after_pre: Tensor1D,
508 attn_norm: Tensor1D,
509 att_x_prev_old: Tensor1D,
510 ffn_x_prev_old: Tensor1D,
511 att_state_old: Tensor1D,
512 xr: Tensor1D,
513 xw: Tensor1D,
514 xk: Tensor1D,
515 xv: Tensor1D,
516 xa: Tensor1D,
517 xg: Tensor1D,
518 r: Tensor1D,
519 k_pre: Tensor1D,
520 k: Tensor1D,
521 v_pre: Tensor1D,
522 v: Tensor1D,
523 nu: Tensor1D,
524 w_hidden: Tensor1D,
525 w_pre: Tensor1D,
526 w_sigmoid: Tensor1D,
527 w_decay: Tensor1D,
528 a_hidden: Tensor1D,
529 a: Tensor1D,
530 g_hidden: Tensor1D,
531 g: Tensor1D,
532 kk_pre: Tensor1D,
533 kk: Tensor1D,
534 y_wkv: Tensor1D,
535 y_gn: Tensor1D,
536 alpha: Tensor1D,
537 y_head: Tensor1D,
538 y_gate: Tensor1D,
539 att_out: Tensor1D,
540 x_after_attn: Tensor1D,
541 ffn_norm: Tensor1D,
542 ffn_xk: Tensor1D,
543 ffn_pre: Tensor1D,
544 ffn_k: Tensor1D,
545 ffn_out: Tensor1D,
546 x_out: Tensor1D,
547 v_hidden: Tensor1D,
548 uses_v_residual: bool,
549}
550
551impl LayerTrainTrace {
552 fn new(cfg: &Config) -> Self {
553 let c = cfg.hidden_size;
554 let i = cfg.intermediate_size;
555 let state = cfg.num_heads * cfg.head_dim * cfg.head_dim;
556 Self {
557 x_in: Tensor1D::zeros(c),
558 x_after_pre: Tensor1D::zeros(c),
559 attn_norm: Tensor1D::zeros(c),
560 att_x_prev_old: Tensor1D::zeros(c),
561 ffn_x_prev_old: Tensor1D::zeros(c),
562 att_state_old: Tensor1D::zeros(state),
563 xr: Tensor1D::zeros(c),
564 xw: Tensor1D::zeros(c),
565 xk: Tensor1D::zeros(c),
566 xv: Tensor1D::zeros(c),
567 xa: Tensor1D::zeros(c),
568 xg: Tensor1D::zeros(c),
569 r: Tensor1D::zeros(c),
570 k_pre: Tensor1D::zeros(c),
571 k: Tensor1D::zeros(c),
572 v_pre: Tensor1D::zeros(c),
573 v: Tensor1D::zeros(c),
574 nu: Tensor1D::zeros(c),
575 w_hidden: Tensor1D::zeros(cfg.decay_low_rank),
576 w_pre: Tensor1D::zeros(c),
577 w_sigmoid: Tensor1D::zeros(c),
578 w_decay: Tensor1D::zeros(c),
579 a_hidden: Tensor1D::zeros(cfg.a_low_rank),
580 a: Tensor1D::zeros(c),
581 g_hidden: Tensor1D::zeros(cfg.g_low_rank),
582 g: Tensor1D::zeros(c),
583 kk_pre: Tensor1D::zeros(c),
584 kk: Tensor1D::zeros(c),
585 y_wkv: Tensor1D::zeros(c),
586 y_gn: Tensor1D::zeros(c),
587 alpha: Tensor1D::zeros(cfg.num_heads),
588 y_head: Tensor1D::zeros(c),
589 y_gate: Tensor1D::zeros(c),
590 att_out: Tensor1D::zeros(c),
591 x_after_attn: Tensor1D::zeros(c),
592 ffn_norm: Tensor1D::zeros(c),
593 ffn_xk: Tensor1D::zeros(c),
594 ffn_pre: Tensor1D::zeros(i),
595 ffn_k: Tensor1D::zeros(i),
596 ffn_out: Tensor1D::zeros(c),
597 x_out: Tensor1D::zeros(c),
598 v_hidden: Tensor1D::zeros(cfg.v_low_rank.max(1)),
599 uses_v_residual: false,
600 }
601 }
602}
603
604#[derive(Clone)]
605struct TokenTrainTrace {
606 token: usize,
607 x: Tensor1D,
608 x_normed: Tensor1D,
609 v_first: Tensor1D,
610 layers: Vec<LayerTrainTrace>,
611}
612
613impl TokenTrainTrace {
614 fn from_scratch(scratch: &ScratchBuffers) -> Self {
615 Self {
616 token: scratch.train_token,
617 x: scratch.x.clone(),
618 x_normed: scratch.x_normed.clone(),
619 v_first: scratch.train_v_first.clone(),
620 layers: scratch.train_trace_layers.clone(),
621 }
622 }
623
624 fn clone_from_scratch(&mut self, scratch: &ScratchBuffers) {
625 self.token = scratch.train_token;
626 self.x.clone_from(&scratch.x);
627 self.x_normed.clone_from(&scratch.x_normed);
628 self.v_first.clone_from(&scratch.train_v_first);
629 self.layers.clone_from(&scratch.train_trace_layers);
630 }
631}
632
633#[derive(Clone)]
634struct LayerRecurrentGradState {
635 att_x_prev: Tensor1D,
636 att_state: Tensor1D,
637 ffn_x_prev: Tensor1D,
638}
639
640impl LayerRecurrentGradState {
641 fn new(cfg: &Config) -> Self {
642 let state_size = cfg.num_heads * cfg.head_dim * cfg.head_dim;
643 Self {
644 att_x_prev: Tensor1D::zeros(cfg.hidden_size),
645 att_state: Tensor1D::zeros(state_size),
646 ffn_x_prev: Tensor1D::zeros(cfg.hidden_size),
647 }
648 }
649}
650
651#[derive(Clone)]
652struct RecurrentGradState {
653 layers: Vec<LayerRecurrentGradState>,
654}
655
656impl RecurrentGradState {
657 fn new(cfg: &Config) -> Self {
658 Self {
659 layers: (0..cfg.num_layers)
660 .map(|_| LayerRecurrentGradState::new(cfg))
661 .collect(),
662 }
663 }
664
665 fn zero(&mut self) {
666 for layer in &mut self.layers {
667 layer.att_x_prev.zero();
668 layer.att_state.zero();
669 layer.ffn_x_prev.zero();
670 }
671 }
672}
673
674#[derive(Clone)]
676pub struct ScratchBuffers {
677 x: Tensor1D, x_normed: Tensor1D, xr: Tensor1D, xw: Tensor1D, xk: Tensor1D, xv: Tensor1D, xa: Tensor1D, xg: Tensor1D, r: Tensor1D, k: Tensor1D, v: Tensor1D, w_lora_tmp: Tensor1D, w_decay: Tensor1D, a: Tensor1D, g: Tensor1D, kk: Tensor1D, y: Tensor1D, att_out: Tensor1D, ffn_k: Tensor1D, ffn_out: Tensor1D, logits: Tensor1D, grad_x: Tensor1D,
699 grad_x2: Tensor1D,
700 grad_x3: Tensor1D,
701 grad_x4: Tensor1D,
702 grad_x5: Tensor1D,
703 grad_x6: Tensor1D,
704 grad_v_first: Tensor1D,
705 grad_param: Tensor1D,
706 grad_param2: Tensor1D,
707 grad_saved: Tensor1D,
708 grad_ffn: Tensor1D,
709 grad_ffn2: Tensor1D,
710 grad_low_rank: Tensor1D,
711 grad_low_rank2: Tensor1D,
712 grad_att_state: Tensor1D,
713 grad_logits: Tensor1D,
714 train_trace_layers: Vec<LayerTrainTrace>,
715 train_token: usize,
716 train_v_first: Tensor1D,
717 train_trace_valid: bool,
718 capture_train_trace: bool,
719}
720
721pub(crate) struct TbpttReplayWorkspace {
722 grads: FullGradState,
723 recurrent: RecurrentGradState,
724 bias_grad: Vec<f32>,
725 checkpoint_state: State,
726 replay_state: State,
727 checkpoints: Vec<State>,
728 step_states: Vec<State>,
729 step_traces: Vec<TokenTrainTrace>,
730 step_pdfs: Vec<f64>,
731}
732
733impl TbpttReplayWorkspace {
734 pub(crate) fn new(model: &Model) -> Self {
735 Self {
736 grads: model.new_full_grad_state(),
737 recurrent: model.new_recurrent_grad_state(),
738 bias_grad: Vec::new(),
739 checkpoint_state: model.new_state(),
740 replay_state: model.new_state(),
741 checkpoints: Vec::new(),
742 step_states: Vec::new(),
743 step_traces: Vec::new(),
744 step_pdfs: Vec::new(),
745 }
746 }
747}
748
749fn ensure_cloned_len<T: Clone>(buf: &mut Vec<T>, len: usize, template: &T) {
750 if buf.len() < len {
751 buf.resize_with(len, || template.clone());
752 }
753}
754
755impl ScratchBuffers {
756 pub fn new(cfg: &Config) -> Self {
758 let c = cfg.hidden_size;
759 let i = cfg.intermediate_size;
760 let v = cfg.vocab_size;
761 let state_size = cfg.num_heads * cfg.head_dim * cfg.head_dim;
762 let d_rank = cfg
763 .decay_low_rank
764 .max(cfg.a_low_rank)
765 .max(cfg.v_low_rank)
766 .max(cfg.g_low_rank)
767 .max(64);
768 let mut train_trace_layers = Vec::with_capacity(cfg.num_layers);
769 for _ in 0..cfg.num_layers {
770 train_trace_layers.push(LayerTrainTrace::new(cfg));
771 }
772
773 Self {
774 x: Tensor1D::zeros(c),
775 x_normed: Tensor1D::zeros(c),
776 xr: Tensor1D::zeros(c),
777 xw: Tensor1D::zeros(c),
778 xk: Tensor1D::zeros(c),
779 xv: Tensor1D::zeros(c),
780 xa: Tensor1D::zeros(c),
781 xg: Tensor1D::zeros(c),
782 r: Tensor1D::zeros(c),
783 k: Tensor1D::zeros(c),
784 v: Tensor1D::zeros(c),
785 w_lora_tmp: Tensor1D::zeros(d_rank),
786 w_decay: Tensor1D::zeros(c),
787 a: Tensor1D::zeros(c),
788 g: Tensor1D::zeros(c),
789 kk: Tensor1D::zeros(c),
790 y: Tensor1D::zeros(c),
791 att_out: Tensor1D::zeros(c),
792 ffn_k: Tensor1D::zeros(i),
793 ffn_out: Tensor1D::zeros(c),
794 logits: Tensor1D::zeros(v),
795 grad_x: Tensor1D::zeros(c),
796 grad_x2: Tensor1D::zeros(c),
797 grad_x3: Tensor1D::zeros(c),
798 grad_x4: Tensor1D::zeros(c),
799 grad_x5: Tensor1D::zeros(c),
800 grad_x6: Tensor1D::zeros(c),
801 grad_v_first: Tensor1D::zeros(c),
802 grad_param: Tensor1D::zeros(c),
803 grad_param2: Tensor1D::zeros(c),
804 grad_saved: Tensor1D::zeros(c),
805 grad_ffn: Tensor1D::zeros(i),
806 grad_ffn2: Tensor1D::zeros(i),
807 grad_low_rank: Tensor1D::zeros(d_rank),
808 grad_low_rank2: Tensor1D::zeros(d_rank),
809 grad_att_state: Tensor1D::zeros(state_size),
810 grad_logits: Tensor1D::zeros(v),
811 train_trace_layers,
812 train_token: 0,
813 train_v_first: Tensor1D::zeros(c),
814 train_trace_valid: false,
815 capture_train_trace: false,
816 }
817 }
818
819 #[inline]
821 pub fn lm_head_input(&self) -> &[f32] {
822 self.x_normed.as_slice()
823 }
824
825 #[inline]
826 pub fn logits(&self) -> &[f32] {
828 self.logits.as_slice()
829 }
830
831 #[inline]
833 pub fn set_lm_head_input(&mut self, value: &[f32]) {
834 self.x_normed.as_mut_slice().copy_from_slice(value);
835 }
836
837 #[inline]
839 pub fn set_capture_train_trace(&mut self, enabled: bool) {
840 self.capture_train_trace = enabled;
841 if !enabled {
842 self.train_trace_valid = false;
843 }
844 }
845
846 #[inline]
848 pub fn has_train_trace(&self) -> bool {
849 self.train_trace_valid
850 }
851}
852
853impl Model {
854 fn tensor_from(weights: &Weights, name: &str) -> Result<Tensor1D> {
855 Ok(Tensor1D::from_vec(weights.require(name)?.data().to_vec()))
856 }
857
858 fn optional_tensor_from(weights: &Weights, name: &str) -> Option<Tensor1D> {
859 weights
860 .get(name)
861 .map(|tensor| Tensor1D::from_vec(tensor.data().to_vec()))
862 }
863
864 pub fn load<P: AsRef<Path>>(path: P) -> Result<Self> {
866 let weights = Weights::load(path.as_ref()).with_context(|| {
867 format!(
868 "Failed to load model weights from {}",
869 path.as_ref().display()
870 )
871 })?;
872
873 let emb = weights.require("model.embeddings.weight")?;
875 let vocab_size = emb.shape()[0];
876 let hidden_size = emb.shape()[1];
877
878 let num_heads = hidden_size / 64; let head_dim = 64;
880
881 let mut num_layers = 0;
883 while weights
884 .get(&format!("model.layers.{}.attn.r_proj.weight", num_layers))
885 .is_some()
886 {
887 num_layers += 1;
888 }
889
890 let ffn_key = weights.require("model.layers.0.ffn.key.weight")?;
892 let intermediate_size = ffn_key.shape()[0];
893
894 let w1 = weights.require("model.layers.0.attn.w_lora.lora.0.weight")?;
896 let decay_low_rank = w1.shape()[0];
897
898 let a1 = weights.require("model.layers.0.attn.a_lora.lora.0.weight")?;
899 let a_low_rank = a1.shape()[0];
900
901 let g1 = weights.require("model.layers.0.attn.g_lora.lora.0.weight")?;
902 let g_low_rank = g1.shape()[0];
903
904 let v_low_rank = if num_layers > 1 {
906 if let Some(v1) = weights.get("model.layers.1.attn.v_lora.lora.0.weight") {
907 v1.shape()[0]
908 } else {
909 32
910 }
911 } else {
912 32
913 };
914
915 let cfg = Config {
916 vocab_size,
917 hidden_size,
918 num_layers,
919 num_heads,
920 head_dim,
921 intermediate_size,
922 layer_norm_eps: 1e-5,
923 group_norm_eps: 64e-5,
924 decay_low_rank,
925 a_low_rank,
926 v_low_rank,
927 g_low_rank,
928 };
929
930 let embeddings = Self::tensor_from(&weights, "model.embeddings.weight")?;
932
933 let ln_out_w = Self::tensor_from(&weights, "model.norm.weight")?;
935 let ln_out_b = Self::tensor_from(&weights, "model.norm.bias")?;
936
937 let lm_head = Self::tensor_from(&weights, "lm_head.weight")?;
939
940 let mut blocks = Vec::with_capacity(num_layers);
942 for i in 0..num_layers {
943 let prefix = format!("model.layers.{}", i);
944
945 let (pre_norm_w, pre_norm_b) = if i == 0 {
947 (
948 Some(Self::tensor_from(
949 &weights,
950 &format!("{}.pre_norm.weight", prefix),
951 )?),
952 Some(Self::tensor_from(
953 &weights,
954 &format!("{}.pre_norm.bias", prefix),
955 )?),
956 )
957 } else {
958 (None, None)
959 };
960
961 let attn_norm_w = Self::tensor_from(&weights, &format!("{}.attn_norm.weight", prefix))?;
963 let attn_norm_b = Self::tensor_from(&weights, &format!("{}.attn_norm.bias", prefix))?;
964 let ffn_norm_w = Self::tensor_from(&weights, &format!("{}.ffn_norm.weight", prefix))?;
965 let ffn_norm_b = Self::tensor_from(&weights, &format!("{}.ffn_norm.bias", prefix))?;
966
967 let r_proj_data = weights
970 .require(&format!("{}.attn.r_proj.weight", prefix))?
971 .data();
972 let k_proj_data = weights
973 .require(&format!("{}.attn.k_proj.weight", prefix))?
974 .data();
975 let v_proj_data = weights
976 .require(&format!("{}.attn.v_proj.weight", prefix))?
977 .data();
978
979 let proj_size = hidden_size * hidden_size;
981 let mut rkv_proj = Tensor1D::zeros(3 * proj_size);
982 rkv_proj.as_mut_slice()[0..proj_size].copy_from_slice(r_proj_data);
983 rkv_proj.as_mut_slice()[proj_size..2 * proj_size].copy_from_slice(k_proj_data);
984 rkv_proj.as_mut_slice()[2 * proj_size..3 * proj_size].copy_from_slice(v_proj_data);
985
986 let attn = AttentionWeights {
987 x_r: Self::tensor_from(&weights, &format!("{}.attn.x_r", prefix))?,
988 x_w: Self::tensor_from(&weights, &format!("{}.attn.x_w", prefix))?,
989 x_k: Self::tensor_from(&weights, &format!("{}.attn.x_k", prefix))?,
990 x_v: Self::tensor_from(&weights, &format!("{}.attn.x_v", prefix))?,
991 x_a: Self::tensor_from(&weights, &format!("{}.attn.x_a", prefix))?,
992 x_g: Self::tensor_from(&weights, &format!("{}.attn.x_g", prefix))?,
993
994 rkv_proj,
995 o_proj: Self::tensor_from(&weights, &format!("{}.attn.o_proj.weight", prefix))?,
996
997 w1: Self::tensor_from(&weights, &format!("{}.attn.w_lora.lora.0.weight", prefix))?,
998 w2: Self::tensor_from(&weights, &format!("{}.attn.w_lora.lora.2.weight", prefix))?,
999 w0: Self::tensor_from(&weights, &format!("{}.attn.w_lora.lora.2.bias", prefix))?,
1000
1001 a1: Self::tensor_from(&weights, &format!("{}.attn.a_lora.lora.0.weight", prefix))?,
1002 a2: Self::tensor_from(&weights, &format!("{}.attn.a_lora.lora.2.weight", prefix))?,
1003 a0: Self::tensor_from(&weights, &format!("{}.attn.a_lora.lora.2.bias", prefix))?,
1004
1005 v1: Self::optional_tensor_from(
1006 &weights,
1007 &format!("{}.attn.v_lora.lora.0.weight", prefix),
1008 ),
1009 v2: Self::optional_tensor_from(
1010 &weights,
1011 &format!("{}.attn.v_lora.lora.2.weight", prefix),
1012 ),
1013 v0: Self::optional_tensor_from(
1014 &weights,
1015 &format!("{}.attn.v_lora.lora.2.bias", prefix),
1016 ),
1017
1018 g1: Self::tensor_from(&weights, &format!("{}.attn.g_lora.lora.0.weight", prefix))?,
1019 g2: Self::tensor_from(&weights, &format!("{}.attn.g_lora.lora.2.weight", prefix))?,
1020
1021 k_k: Self::tensor_from(&weights, &format!("{}.attn.k_k", prefix))?,
1022 k_a: Self::tensor_from(&weights, &format!("{}.attn.k_a", prefix))?,
1023 r_k: Self::tensor_from(&weights, &format!("{}.attn.r_k", prefix))?,
1024
1025 g_norm_w: Self::tensor_from(&weights, &format!("{}.attn.g_norm.weight", prefix))?,
1026 g_norm_b: Self::tensor_from(&weights, &format!("{}.attn.g_norm.bias", prefix))?,
1027 };
1028
1029 let ffn = FfnWeights {
1031 x_k: Self::tensor_from(&weights, &format!("{}.ffn.x_k", prefix))?,
1032 key_w: Self::tensor_from(&weights, &format!("{}.ffn.key.weight", prefix))?,
1033 value_w: Self::tensor_from(&weights, &format!("{}.ffn.value.weight", prefix))?,
1034 };
1035
1036 blocks.push(BlockWeights {
1037 pre_norm_w,
1038 pre_norm_b,
1039 attn_norm_w,
1040 attn_norm_b,
1041 ffn_norm_w,
1042 ffn_norm_b,
1043 attn,
1044 ffn,
1045 });
1046 }
1047
1048 Ok(Self {
1049 cfg,
1050 embeddings,
1051 ln_out_w,
1052 ln_out_b,
1053 lm_head,
1054 blocks,
1055 })
1056 }
1057
1058 pub fn new_random(cfg: Config, seed: u64) -> Result<Self> {
1060 cfg.validate()?;
1061
1062 let mut rng = RwkvRng::new(seed);
1063 let c = cfg.hidden_size;
1064 let v = cfg.vocab_size;
1065 let i = cfg.intermediate_size;
1066 let d_w = cfg.decay_low_rank;
1067 let d_a = cfg.a_low_rank;
1068 let d_v = cfg.v_low_rank;
1069 let d_g = cfg.g_low_rank;
1070
1071 let mut embeddings = Tensor1D::zeros(v * c);
1072 init_uniform(&mut embeddings, &mut rng, 0.02);
1073
1074 let mut ln_out_w = Tensor1D::zeros(c);
1075 let mut ln_out_b = Tensor1D::zeros(c);
1076 init_const(&mut ln_out_w, 1.0);
1077 init_const(&mut ln_out_b, 0.0);
1078
1079 let mut lm_head = Tensor1D::zeros(v * c);
1080 init_uniform(&mut lm_head, &mut rng, 0.02);
1081
1082 let mut blocks = Vec::with_capacity(cfg.num_layers);
1083 for layer_idx in 0..cfg.num_layers {
1084 let (pre_norm_w, pre_norm_b) = if layer_idx == 0 {
1085 let mut w = Tensor1D::zeros(c);
1086 let mut b = Tensor1D::zeros(c);
1087 init_const(&mut w, 1.0);
1088 init_const(&mut b, 0.0);
1089 (Some(w), Some(b))
1090 } else {
1091 (None, None)
1092 };
1093
1094 let mut attn_norm_w = Tensor1D::zeros(c);
1095 let mut attn_norm_b = Tensor1D::zeros(c);
1096 init_const(&mut attn_norm_w, 1.0);
1097 init_const(&mut attn_norm_b, 0.0);
1098
1099 let mut ffn_norm_w = Tensor1D::zeros(c);
1100 let mut ffn_norm_b = Tensor1D::zeros(c);
1101 init_const(&mut ffn_norm_w, 1.0);
1102 init_const(&mut ffn_norm_b, 0.0);
1103
1104 let mut rkv_proj = Tensor1D::zeros(3 * c * c);
1105 init_uniform(&mut rkv_proj, &mut rng, 0.02);
1106
1107 let mut o_proj = Tensor1D::zeros(c * c);
1108 init_uniform(&mut o_proj, &mut rng, 0.02);
1109
1110 let mut w1 = Tensor1D::zeros(d_w * c);
1111 let mut w2 = Tensor1D::zeros(c * d_w);
1112 let mut w0 = Tensor1D::zeros(c);
1113 init_uniform(&mut w1, &mut rng, 0.02);
1114 init_uniform(&mut w2, &mut rng, 0.02);
1115 init_const(&mut w0, 0.0);
1116
1117 let mut a1 = Tensor1D::zeros(d_a * c);
1118 let mut a2 = Tensor1D::zeros(c * d_a);
1119 let mut a0 = Tensor1D::zeros(c);
1120 init_uniform(&mut a1, &mut rng, 0.02);
1121 init_uniform(&mut a2, &mut rng, 0.02);
1122 init_const(&mut a0, 0.0);
1123
1124 let (v1, v2, v0) = if layer_idx == 0 {
1125 (None, None, None)
1126 } else {
1127 let mut v1 = Tensor1D::zeros(d_v * c);
1128 let mut v2 = Tensor1D::zeros(c * d_v);
1129 let mut v0 = Tensor1D::zeros(c);
1130 init_uniform(&mut v1, &mut rng, 0.02);
1131 init_uniform(&mut v2, &mut rng, 0.02);
1132 init_const(&mut v0, 0.0);
1133 (Some(v1), Some(v2), Some(v0))
1134 };
1135
1136 let mut g1 = Tensor1D::zeros(d_g * c);
1137 let mut g2 = Tensor1D::zeros(c * d_g);
1138 init_uniform(&mut g1, &mut rng, 0.02);
1139 init_uniform(&mut g2, &mut rng, 0.02);
1140
1141 let mut x_r = Tensor1D::zeros(c);
1142 let mut x_w = Tensor1D::zeros(c);
1143 let mut x_k = Tensor1D::zeros(c);
1144 let mut x_v = Tensor1D::zeros(c);
1145 let mut x_a = Tensor1D::zeros(c);
1146 let mut x_g = Tensor1D::zeros(c);
1147 init_centered(&mut x_r, &mut rng, 0.5, 0.02);
1148 init_centered(&mut x_w, &mut rng, 0.5, 0.02);
1149 init_centered(&mut x_k, &mut rng, 0.5, 0.02);
1150 init_centered(&mut x_v, &mut rng, 0.5, 0.02);
1151 init_centered(&mut x_a, &mut rng, 0.5, 0.02);
1152 init_centered(&mut x_g, &mut rng, 0.5, 0.02);
1153
1154 let mut k_k = Tensor1D::zeros(c);
1155 let mut k_a = Tensor1D::zeros(c);
1156 let mut r_k = Tensor1D::zeros(c);
1157 init_const(&mut k_k, 1.0);
1158 init_const(&mut k_a, 1.0);
1159 init_const(&mut r_k, 1.0);
1160
1161 let mut g_norm_w = Tensor1D::zeros(c);
1162 let mut g_norm_b = Tensor1D::zeros(c);
1163 init_const(&mut g_norm_w, 1.0);
1164 init_const(&mut g_norm_b, 0.0);
1165
1166 let attn = AttentionWeights {
1167 x_r,
1168 x_w,
1169 x_k,
1170 x_v,
1171 x_a,
1172 x_g,
1173 rkv_proj,
1174 o_proj,
1175 w1,
1176 w2,
1177 w0,
1178 a1,
1179 a2,
1180 a0,
1181 v1,
1182 v2,
1183 v0,
1184 g1,
1185 g2,
1186 k_k,
1187 k_a,
1188 r_k,
1189 g_norm_w,
1190 g_norm_b,
1191 };
1192
1193 let mut ffn_x_k = Tensor1D::zeros(c);
1194 init_centered(&mut ffn_x_k, &mut rng, 0.5, 0.02);
1195 let mut key_w = Tensor1D::zeros(i * c);
1196 let mut value_w = Tensor1D::zeros(c * i);
1197 init_uniform(&mut key_w, &mut rng, 0.02);
1198 init_uniform(&mut value_w, &mut rng, 0.02);
1199
1200 let ffn = FfnWeights {
1201 x_k: ffn_x_k,
1202 key_w,
1203 value_w,
1204 };
1205
1206 blocks.push(BlockWeights {
1207 pre_norm_w,
1208 pre_norm_b,
1209 attn_norm_w,
1210 attn_norm_b,
1211 ffn_norm_w,
1212 ffn_norm_b,
1213 attn,
1214 ffn,
1215 });
1216 }
1217
1218 Ok(Self {
1219 cfg,
1220 embeddings,
1221 ln_out_w,
1222 ln_out_b,
1223 lm_head,
1224 blocks,
1225 })
1226 }
1227
1228 pub fn save_safetensors<P: AsRef<Path>>(&self, path: P) -> Result<()> {
1230 #[derive(Clone)]
1231 struct TensorRec {
1232 name: String,
1233 shape: Vec<usize>,
1234 data: Vec<f32>,
1235 }
1236
1237 let c = self.cfg.hidden_size;
1238 let v = self.cfg.vocab_size;
1239 let i = self.cfg.intermediate_size;
1240 let d_w = self.cfg.decay_low_rank;
1241 let d_a = self.cfg.a_low_rank;
1242 let d_v = self.cfg.v_low_rank;
1243 let d_g = self.cfg.g_low_rank;
1244
1245 let mut recs = Vec::<TensorRec>::new();
1246 let push = |recs: &mut Vec<TensorRec>, name: String, shape: Vec<usize>, src: &Tensor1D| {
1247 recs.push(TensorRec {
1248 name,
1249 shape,
1250 data: src.as_slice().to_vec(),
1251 });
1252 };
1253
1254 push(
1255 &mut recs,
1256 "model.embeddings.weight".to_string(),
1257 vec![v, c],
1258 &self.embeddings,
1259 );
1260 push(
1261 &mut recs,
1262 "model.norm.weight".to_string(),
1263 vec![c],
1264 &self.ln_out_w,
1265 );
1266 push(
1267 &mut recs,
1268 "model.norm.bias".to_string(),
1269 vec![c],
1270 &self.ln_out_b,
1271 );
1272 push(
1273 &mut recs,
1274 "lm_head.weight".to_string(),
1275 vec![v, c],
1276 &self.lm_head,
1277 );
1278
1279 for (idx, b) in self.blocks.iter().enumerate() {
1280 let pfx = format!("model.layers.{idx}");
1281 if let (Some(w), Some(bias)) = (&b.pre_norm_w, &b.pre_norm_b) {
1282 push(&mut recs, format!("{pfx}.pre_norm.weight"), vec![c], w);
1283 push(&mut recs, format!("{pfx}.pre_norm.bias"), vec![c], bias);
1284 }
1285
1286 push(
1287 &mut recs,
1288 format!("{pfx}.attn_norm.weight"),
1289 vec![c],
1290 &b.attn_norm_w,
1291 );
1292 push(
1293 &mut recs,
1294 format!("{pfx}.attn_norm.bias"),
1295 vec![c],
1296 &b.attn_norm_b,
1297 );
1298 push(
1299 &mut recs,
1300 format!("{pfx}.ffn_norm.weight"),
1301 vec![c],
1302 &b.ffn_norm_w,
1303 );
1304 push(
1305 &mut recs,
1306 format!("{pfx}.ffn_norm.bias"),
1307 vec![c],
1308 &b.ffn_norm_b,
1309 );
1310
1311 let proj = b.attn.rkv_proj.as_slice();
1312 let proj_size = c * c;
1313 recs.push(TensorRec {
1314 name: format!("{pfx}.attn.r_proj.weight"),
1315 shape: vec![c, c],
1316 data: proj[0..proj_size].to_vec(),
1317 });
1318 recs.push(TensorRec {
1319 name: format!("{pfx}.attn.k_proj.weight"),
1320 shape: vec![c, c],
1321 data: proj[proj_size..2 * proj_size].to_vec(),
1322 });
1323 recs.push(TensorRec {
1324 name: format!("{pfx}.attn.v_proj.weight"),
1325 shape: vec![c, c],
1326 data: proj[2 * proj_size..3 * proj_size].to_vec(),
1327 });
1328
1329 push(
1330 &mut recs,
1331 format!("{pfx}.attn.o_proj.weight"),
1332 vec![c, c],
1333 &b.attn.o_proj,
1334 );
1335 push(&mut recs, format!("{pfx}.attn.x_r"), vec![c], &b.attn.x_r);
1336 push(&mut recs, format!("{pfx}.attn.x_w"), vec![c], &b.attn.x_w);
1337 push(&mut recs, format!("{pfx}.attn.x_k"), vec![c], &b.attn.x_k);
1338 push(&mut recs, format!("{pfx}.attn.x_v"), vec![c], &b.attn.x_v);
1339 push(&mut recs, format!("{pfx}.attn.x_a"), vec![c], &b.attn.x_a);
1340 push(&mut recs, format!("{pfx}.attn.x_g"), vec![c], &b.attn.x_g);
1341
1342 push(
1343 &mut recs,
1344 format!("{pfx}.attn.w_lora.lora.0.weight"),
1345 vec![d_w, c],
1346 &b.attn.w1,
1347 );
1348 push(
1349 &mut recs,
1350 format!("{pfx}.attn.w_lora.lora.2.weight"),
1351 vec![c, d_w],
1352 &b.attn.w2,
1353 );
1354 push(
1355 &mut recs,
1356 format!("{pfx}.attn.w_lora.lora.2.bias"),
1357 vec![c],
1358 &b.attn.w0,
1359 );
1360
1361 push(
1362 &mut recs,
1363 format!("{pfx}.attn.a_lora.lora.0.weight"),
1364 vec![d_a, c],
1365 &b.attn.a1,
1366 );
1367 push(
1368 &mut recs,
1369 format!("{pfx}.attn.a_lora.lora.2.weight"),
1370 vec![c, d_a],
1371 &b.attn.a2,
1372 );
1373 push(
1374 &mut recs,
1375 format!("{pfx}.attn.a_lora.lora.2.bias"),
1376 vec![c],
1377 &b.attn.a0,
1378 );
1379
1380 if let Some(v1) = &b.attn.v1 {
1381 push(
1382 &mut recs,
1383 format!("{pfx}.attn.v_lora.lora.0.weight"),
1384 vec![d_v, c],
1385 v1,
1386 );
1387 }
1388 if let Some(v2) = &b.attn.v2 {
1389 push(
1390 &mut recs,
1391 format!("{pfx}.attn.v_lora.lora.2.weight"),
1392 vec![c, d_v],
1393 v2,
1394 );
1395 }
1396 if let Some(v0) = &b.attn.v0 {
1397 push(
1398 &mut recs,
1399 format!("{pfx}.attn.v_lora.lora.2.bias"),
1400 vec![c],
1401 v0,
1402 );
1403 }
1404
1405 push(
1406 &mut recs,
1407 format!("{pfx}.attn.g_lora.lora.0.weight"),
1408 vec![d_g, c],
1409 &b.attn.g1,
1410 );
1411 push(
1412 &mut recs,
1413 format!("{pfx}.attn.g_lora.lora.2.weight"),
1414 vec![c, d_g],
1415 &b.attn.g2,
1416 );
1417
1418 push(&mut recs, format!("{pfx}.attn.k_k"), vec![c], &b.attn.k_k);
1419 push(&mut recs, format!("{pfx}.attn.k_a"), vec![c], &b.attn.k_a);
1420 push(&mut recs, format!("{pfx}.attn.r_k"), vec![c], &b.attn.r_k);
1421 push(
1422 &mut recs,
1423 format!("{pfx}.attn.g_norm.weight"),
1424 vec![c],
1425 &b.attn.g_norm_w,
1426 );
1427 push(
1428 &mut recs,
1429 format!("{pfx}.attn.g_norm.bias"),
1430 vec![c],
1431 &b.attn.g_norm_b,
1432 );
1433
1434 push(&mut recs, format!("{pfx}.ffn.x_k"), vec![c], &b.ffn.x_k);
1435 push(
1436 &mut recs,
1437 format!("{pfx}.ffn.key.weight"),
1438 vec![i, c],
1439 &b.ffn.key_w,
1440 );
1441 push(
1442 &mut recs,
1443 format!("{pfx}.ffn.value.weight"),
1444 vec![c, i],
1445 &b.ffn.value_w,
1446 );
1447 }
1448
1449 recs.sort_by(|a, b| a.name.cmp(&b.name));
1450 let mut offset = 0usize;
1451 let mut header = serde_json::Map::new();
1452 header.insert("__metadata__".to_string(), json!({}));
1453 for rec in &recs {
1454 let bytes = rec.data.len() * 4;
1455 header.insert(
1456 rec.name.clone(),
1457 json!({
1458 "dtype": "F32",
1459 "shape": rec.shape,
1460 "data_offsets": [offset, offset + bytes]
1461 }),
1462 );
1463 offset += bytes;
1464 }
1465 let header_bytes = serde_json::to_vec(&header)?;
1466 let mut f = File::create(path.as_ref())?;
1467 f.write_all(&(header_bytes.len() as u64).to_le_bytes())?;
1468 f.write_all(&header_bytes)?;
1469 for rec in &recs {
1470 for v in &rec.data {
1471 f.write_all(&v.to_le_bytes())?;
1472 }
1473 }
1474 Ok(())
1475 }
1476
1477 pub fn new_full_adam_state(&self) -> FullAdamState {
1479 let mut blocks = Vec::with_capacity(self.blocks.len());
1480 for b in &self.blocks {
1481 blocks.push(BlockAdamState {
1482 pre_norm_w: b.pre_norm_w.as_ref().map(|t| AdamTensorState::new(t.len())),
1483 pre_norm_b: b.pre_norm_b.as_ref().map(|t| AdamTensorState::new(t.len())),
1484 attn_norm_w: AdamTensorState::new(b.attn_norm_w.len()),
1485 attn_norm_b: AdamTensorState::new(b.attn_norm_b.len()),
1486 ffn_norm_w: AdamTensorState::new(b.ffn_norm_w.len()),
1487 ffn_norm_b: AdamTensorState::new(b.ffn_norm_b.len()),
1488 attn: AttentionAdamState {
1489 x_r: AdamTensorState::new(b.attn.x_r.len()),
1490 x_w: AdamTensorState::new(b.attn.x_w.len()),
1491 x_k: AdamTensorState::new(b.attn.x_k.len()),
1492 x_v: AdamTensorState::new(b.attn.x_v.len()),
1493 x_a: AdamTensorState::new(b.attn.x_a.len()),
1494 x_g: AdamTensorState::new(b.attn.x_g.len()),
1495 rkv_proj: AdamTensorState::new(b.attn.rkv_proj.len()),
1496 o_proj: AdamTensorState::new(b.attn.o_proj.len()),
1497 w1: AdamTensorState::new(b.attn.w1.len()),
1498 w2: AdamTensorState::new(b.attn.w2.len()),
1499 w0: AdamTensorState::new(b.attn.w0.len()),
1500 a1: AdamTensorState::new(b.attn.a1.len()),
1501 a2: AdamTensorState::new(b.attn.a2.len()),
1502 a0: AdamTensorState::new(b.attn.a0.len()),
1503 v1: b.attn.v1.as_ref().map(|t| AdamTensorState::new(t.len())),
1504 v2: b.attn.v2.as_ref().map(|t| AdamTensorState::new(t.len())),
1505 v0: b.attn.v0.as_ref().map(|t| AdamTensorState::new(t.len())),
1506 g1: AdamTensorState::new(b.attn.g1.len()),
1507 g2: AdamTensorState::new(b.attn.g2.len()),
1508 k_k: AdamTensorState::new(b.attn.k_k.len()),
1509 k_a: AdamTensorState::new(b.attn.k_a.len()),
1510 r_k: AdamTensorState::new(b.attn.r_k.len()),
1511 g_norm_w: AdamTensorState::new(b.attn.g_norm_w.len()),
1512 g_norm_b: AdamTensorState::new(b.attn.g_norm_b.len()),
1513 },
1514 ffn: FfnAdamState {
1515 x_k: AdamTensorState::new(b.ffn.x_k.len()),
1516 key_w: AdamTensorState::new(b.ffn.key_w.len()),
1517 value_w: AdamTensorState::new(b.ffn.value_w.len()),
1518 },
1519 });
1520 }
1521 FullAdamState {
1522 embeddings: AdamTensorState::new(self.embeddings.len()),
1523 ln_out_w: AdamTensorState::new(self.ln_out_w.len()),
1524 ln_out_b: AdamTensorState::new(self.ln_out_b.len()),
1525 lm_head: AdamTensorState::new(self.lm_head.len()),
1526 blocks,
1527 }
1528 }
1529
1530 fn new_full_grad_state(&self) -> FullGradState {
1532 let mut blocks = Vec::with_capacity(self.blocks.len());
1533 for b in &self.blocks {
1534 blocks.push(BlockGradState {
1535 pre_norm_w: b.pre_norm_w.as_ref().map(|t| Tensor1D::zeros(t.len())),
1536 pre_norm_b: b.pre_norm_b.as_ref().map(|t| Tensor1D::zeros(t.len())),
1537 attn_norm_w: Tensor1D::zeros(b.attn_norm_w.len()),
1538 attn_norm_b: Tensor1D::zeros(b.attn_norm_b.len()),
1539 ffn_norm_w: Tensor1D::zeros(b.ffn_norm_w.len()),
1540 ffn_norm_b: Tensor1D::zeros(b.ffn_norm_b.len()),
1541 attn: AttentionGradState {
1542 x_r: Tensor1D::zeros(b.attn.x_r.len()),
1543 x_w: Tensor1D::zeros(b.attn.x_w.len()),
1544 x_k: Tensor1D::zeros(b.attn.x_k.len()),
1545 x_v: Tensor1D::zeros(b.attn.x_v.len()),
1546 x_a: Tensor1D::zeros(b.attn.x_a.len()),
1547 x_g: Tensor1D::zeros(b.attn.x_g.len()),
1548 rkv_proj: Tensor1D::zeros(b.attn.rkv_proj.len()),
1549 o_proj: Tensor1D::zeros(b.attn.o_proj.len()),
1550 w1: Tensor1D::zeros(b.attn.w1.len()),
1551 w2: Tensor1D::zeros(b.attn.w2.len()),
1552 w0: Tensor1D::zeros(b.attn.w0.len()),
1553 a1: Tensor1D::zeros(b.attn.a1.len()),
1554 a2: Tensor1D::zeros(b.attn.a2.len()),
1555 a0: Tensor1D::zeros(b.attn.a0.len()),
1556 v1: b.attn.v1.as_ref().map(|t| Tensor1D::zeros(t.len())),
1557 v2: b.attn.v2.as_ref().map(|t| Tensor1D::zeros(t.len())),
1558 v0: b.attn.v0.as_ref().map(|t| Tensor1D::zeros(t.len())),
1559 g1: Tensor1D::zeros(b.attn.g1.len()),
1560 g2: Tensor1D::zeros(b.attn.g2.len()),
1561 k_k: Tensor1D::zeros(b.attn.k_k.len()),
1562 k_a: Tensor1D::zeros(b.attn.k_a.len()),
1563 r_k: Tensor1D::zeros(b.attn.r_k.len()),
1564 g_norm_w: Tensor1D::zeros(b.attn.g_norm_w.len()),
1565 g_norm_b: Tensor1D::zeros(b.attn.g_norm_b.len()),
1566 },
1567 ffn: FfnGradState {
1568 x_k: Tensor1D::zeros(b.ffn.x_k.len()),
1569 key_w: Tensor1D::zeros(b.ffn.key_w.len()),
1570 value_w: Tensor1D::zeros(b.ffn.value_w.len()),
1571 },
1572 });
1573 }
1574 FullGradState {
1575 embeddings: Tensor1D::zeros(self.embeddings.len()),
1576 ln_out_w: Tensor1D::zeros(self.ln_out_w.len()),
1577 ln_out_b: Tensor1D::zeros(self.ln_out_b.len()),
1578 lm_head: Tensor1D::zeros(self.lm_head.len()),
1579 blocks,
1580 }
1581 }
1582
1583 fn new_recurrent_grad_state(&self) -> RecurrentGradState {
1584 RecurrentGradState::new(&self.cfg)
1585 }
1586
1587 pub fn save_full_adam_safetensors<P: AsRef<Path>>(
1589 &self,
1590 adam: &FullAdamState,
1591 path: P,
1592 ) -> Result<()> {
1593 #[derive(Clone)]
1594 struct TensorRec {
1595 name: String,
1596 shape: Vec<usize>,
1597 data: Vec<f32>,
1598 }
1599 let c = self.cfg.hidden_size;
1600 let i = self.cfg.intermediate_size;
1601 let v = self.cfg.vocab_size;
1602 let h = self.cfg.num_heads;
1603 let n = self.cfg.head_dim;
1604 let d_w = self.cfg.decay_low_rank;
1605 let d_a = self.cfg.a_low_rank;
1606 let d_v = self.cfg.v_low_rank;
1607 let d_g = self.cfg.g_low_rank;
1608 let mut recs = Vec::<TensorRec>::new();
1609 let mut push_state = |name: &str, shape: Vec<usize>, st: &AdamTensorState| {
1610 recs.push(TensorRec {
1611 name: format!("{name}.m"),
1612 shape: shape.clone(),
1613 data: st.m.as_slice().to_vec(),
1614 });
1615 recs.push(TensorRec {
1616 name: format!("{name}.v"),
1617 shape,
1618 data: st.v.as_slice().to_vec(),
1619 });
1620 };
1621
1622 push_state("opt.model.embeddings.weight", vec![v, c], &adam.embeddings);
1623 push_state("opt.model.norm.weight", vec![c], &adam.ln_out_w);
1624 push_state("opt.model.norm.bias", vec![c], &adam.ln_out_b);
1625 push_state("opt.lm_head.weight", vec![v, c], &adam.lm_head);
1626 for (idx, b) in adam.blocks.iter().enumerate() {
1627 let p = format!("opt.model.layers.{idx}");
1628 if let Some(st) = &b.pre_norm_w {
1629 push_state(&format!("{p}.pre_norm.weight"), vec![c], st);
1630 }
1631 if let Some(st) = &b.pre_norm_b {
1632 push_state(&format!("{p}.pre_norm.bias"), vec![c], st);
1633 }
1634 push_state(&format!("{p}.attn_norm.weight"), vec![c], &b.attn_norm_w);
1635 push_state(&format!("{p}.attn_norm.bias"), vec![c], &b.attn_norm_b);
1636 push_state(&format!("{p}.ffn_norm.weight"), vec![c], &b.ffn_norm_w);
1637 push_state(&format!("{p}.ffn_norm.bias"), vec![c], &b.ffn_norm_b);
1638
1639 push_state(&format!("{p}.attn.x_r"), vec![c], &b.attn.x_r);
1640 push_state(&format!("{p}.attn.x_w"), vec![c], &b.attn.x_w);
1641 push_state(&format!("{p}.attn.x_k"), vec![c], &b.attn.x_k);
1642 push_state(&format!("{p}.attn.x_v"), vec![c], &b.attn.x_v);
1643 push_state(&format!("{p}.attn.x_a"), vec![c], &b.attn.x_a);
1644 push_state(&format!("{p}.attn.x_g"), vec![c], &b.attn.x_g);
1645 push_state(
1646 &format!("{p}.attn.rkv_proj"),
1647 vec![3, c, c],
1648 &b.attn.rkv_proj,
1649 );
1650 push_state(
1651 &format!("{p}.attn.o_proj.weight"),
1652 vec![c, c],
1653 &b.attn.o_proj,
1654 );
1655 push_state(
1656 &format!("{p}.attn.w_lora.lora.0.weight"),
1657 vec![d_w, c],
1658 &b.attn.w1,
1659 );
1660 push_state(
1661 &format!("{p}.attn.w_lora.lora.2.weight"),
1662 vec![c, d_w],
1663 &b.attn.w2,
1664 );
1665 push_state(&format!("{p}.attn.w_lora.lora.2.bias"), vec![c], &b.attn.w0);
1666 push_state(
1667 &format!("{p}.attn.a_lora.lora.0.weight"),
1668 vec![d_a, c],
1669 &b.attn.a1,
1670 );
1671 push_state(
1672 &format!("{p}.attn.a_lora.lora.2.weight"),
1673 vec![c, d_a],
1674 &b.attn.a2,
1675 );
1676 push_state(&format!("{p}.attn.a_lora.lora.2.bias"), vec![c], &b.attn.a0);
1677 if let Some(st) = &b.attn.v1 {
1678 push_state(&format!("{p}.attn.v_lora.lora.0.weight"), vec![d_v, c], st);
1679 }
1680 if let Some(st) = &b.attn.v2 {
1681 push_state(&format!("{p}.attn.v_lora.lora.2.weight"), vec![c, d_v], st);
1682 }
1683 if let Some(st) = &b.attn.v0 {
1684 push_state(&format!("{p}.attn.v_lora.lora.2.bias"), vec![c], st);
1685 }
1686 push_state(
1687 &format!("{p}.attn.g_lora.lora.0.weight"),
1688 vec![d_g, c],
1689 &b.attn.g1,
1690 );
1691 push_state(
1692 &format!("{p}.attn.g_lora.lora.2.weight"),
1693 vec![c, d_g],
1694 &b.attn.g2,
1695 );
1696 push_state(&format!("{p}.attn.k_k"), vec![c], &b.attn.k_k);
1697 push_state(&format!("{p}.attn.k_a"), vec![c], &b.attn.k_a);
1698 push_state(&format!("{p}.attn.r_k"), vec![h, n], &b.attn.r_k);
1699 push_state(
1700 &format!("{p}.attn.g_norm.weight"),
1701 vec![c],
1702 &b.attn.g_norm_w,
1703 );
1704 push_state(&format!("{p}.attn.g_norm.bias"), vec![c], &b.attn.g_norm_b);
1705
1706 push_state(&format!("{p}.ffn.x_k"), vec![c], &b.ffn.x_k);
1707 push_state(&format!("{p}.ffn.key.weight"), vec![i, c], &b.ffn.key_w);
1708 push_state(&format!("{p}.ffn.value.weight"), vec![c, i], &b.ffn.value_w);
1709 }
1710
1711 recs.sort_by(|a, b| a.name.cmp(&b.name));
1712 let mut offset = 0usize;
1713 let mut header = serde_json::Map::new();
1714 header.insert("__metadata__".to_string(), json!({}));
1715 for rec in &recs {
1716 let bytes = rec.data.len() * 4;
1717 header.insert(
1718 rec.name.clone(),
1719 json!({
1720 "dtype": "F32",
1721 "shape": rec.shape,
1722 "data_offsets": [offset, offset + bytes],
1723 }),
1724 );
1725 offset += bytes;
1726 }
1727
1728 let header_bytes = serde_json::to_vec(&header)?;
1729 let mut f = File::create(path)?;
1730 f.write_all(&(header_bytes.len() as u64).to_le_bytes())?;
1731 f.write_all(&header_bytes)?;
1732 for rec in &recs {
1733 for v in &rec.data {
1734 f.write_all(&v.to_le_bytes())?;
1735 }
1736 }
1737 Ok(())
1738 }
1739
1740 pub fn load_full_adam_safetensors<P: AsRef<Path>>(&self, path: P) -> Result<FullAdamState> {
1742 let weights = Weights::load(path.as_ref()).with_context(|| {
1743 format!(
1744 "failed to load optimizer moments from {}",
1745 path.as_ref().display()
1746 )
1747 })?;
1748 let mut adam = self.new_full_adam_state();
1749 let load_state = |name: &str, st: &mut AdamTensorState| -> Result<()> {
1750 let m_name = format!("{name}.m");
1751 let v_name = format!("{name}.v");
1752 let m_t = weights
1753 .require(&m_name)
1754 .with_context(|| format!("missing optimizer tensor '{m_name}'"))?;
1755 let v_t = weights
1756 .require(&v_name)
1757 .with_context(|| format!("missing optimizer tensor '{v_name}'"))?;
1758 if m_t.data().len() != st.m.len() {
1759 bail!(
1760 "optimizer tensor '{}' len {} != expected {}",
1761 m_name,
1762 m_t.data().len(),
1763 st.m.len()
1764 );
1765 }
1766 if v_t.data().len() != st.v.len() {
1767 bail!(
1768 "optimizer tensor '{}' len {} != expected {}",
1769 v_name,
1770 v_t.data().len(),
1771 st.v.len()
1772 );
1773 }
1774 st.m.as_mut_slice().copy_from_slice(m_t.data());
1775 st.v.as_mut_slice().copy_from_slice(v_t.data());
1776 Ok(())
1777 };
1778
1779 let c = self.cfg.hidden_size;
1780 let i = self.cfg.intermediate_size;
1781 let v = self.cfg.vocab_size;
1782 let h = self.cfg.num_heads;
1783 let n = self.cfg.head_dim;
1784 let _ = (c, i, v, h, n);
1785 load_state("opt.model.embeddings.weight", &mut adam.embeddings)?;
1786 load_state("opt.model.norm.weight", &mut adam.ln_out_w)?;
1787 load_state("opt.model.norm.bias", &mut adam.ln_out_b)?;
1788 load_state("opt.lm_head.weight", &mut adam.lm_head)?;
1789 for (idx, b) in adam.blocks.iter_mut().enumerate() {
1790 let p = format!("opt.model.layers.{idx}");
1791 if let Some(st) = b.pre_norm_w.as_mut() {
1792 load_state(&format!("{p}.pre_norm.weight"), st)?;
1793 }
1794 if let Some(st) = b.pre_norm_b.as_mut() {
1795 load_state(&format!("{p}.pre_norm.bias"), st)?;
1796 }
1797 load_state(&format!("{p}.attn_norm.weight"), &mut b.attn_norm_w)?;
1798 load_state(&format!("{p}.attn_norm.bias"), &mut b.attn_norm_b)?;
1799 load_state(&format!("{p}.ffn_norm.weight"), &mut b.ffn_norm_w)?;
1800 load_state(&format!("{p}.ffn_norm.bias"), &mut b.ffn_norm_b)?;
1801 load_state(&format!("{p}.attn.x_r"), &mut b.attn.x_r)?;
1802 load_state(&format!("{p}.attn.x_w"), &mut b.attn.x_w)?;
1803 load_state(&format!("{p}.attn.x_k"), &mut b.attn.x_k)?;
1804 load_state(&format!("{p}.attn.x_v"), &mut b.attn.x_v)?;
1805 load_state(&format!("{p}.attn.x_a"), &mut b.attn.x_a)?;
1806 load_state(&format!("{p}.attn.x_g"), &mut b.attn.x_g)?;
1807 load_state(&format!("{p}.attn.rkv_proj"), &mut b.attn.rkv_proj)?;
1808 load_state(&format!("{p}.attn.o_proj.weight"), &mut b.attn.o_proj)?;
1809 load_state(&format!("{p}.attn.w_lora.lora.0.weight"), &mut b.attn.w1)?;
1810 load_state(&format!("{p}.attn.w_lora.lora.2.weight"), &mut b.attn.w2)?;
1811 load_state(&format!("{p}.attn.w_lora.lora.2.bias"), &mut b.attn.w0)?;
1812 load_state(&format!("{p}.attn.a_lora.lora.0.weight"), &mut b.attn.a1)?;
1813 load_state(&format!("{p}.attn.a_lora.lora.2.weight"), &mut b.attn.a2)?;
1814 load_state(&format!("{p}.attn.a_lora.lora.2.bias"), &mut b.attn.a0)?;
1815 if let Some(st) = b.attn.v1.as_mut() {
1816 load_state(&format!("{p}.attn.v_lora.lora.0.weight"), st)?;
1817 }
1818 if let Some(st) = b.attn.v2.as_mut() {
1819 load_state(&format!("{p}.attn.v_lora.lora.2.weight"), st)?;
1820 }
1821 if let Some(st) = b.attn.v0.as_mut() {
1822 load_state(&format!("{p}.attn.v_lora.lora.2.bias"), st)?;
1823 }
1824 load_state(&format!("{p}.attn.g_lora.lora.0.weight"), &mut b.attn.g1)?;
1825 load_state(&format!("{p}.attn.g_lora.lora.2.weight"), &mut b.attn.g2)?;
1826 load_state(&format!("{p}.attn.k_k"), &mut b.attn.k_k)?;
1827 load_state(&format!("{p}.attn.k_a"), &mut b.attn.k_a)?;
1828 load_state(&format!("{p}.attn.r_k"), &mut b.attn.r_k)?;
1829 load_state(&format!("{p}.attn.g_norm.weight"), &mut b.attn.g_norm_w)?;
1830 load_state(&format!("{p}.attn.g_norm.bias"), &mut b.attn.g_norm_b)?;
1831 load_state(&format!("{p}.ffn.x_k"), &mut b.ffn.x_k)?;
1832 load_state(&format!("{p}.ffn.key.weight"), &mut b.ffn.key_w)?;
1833 load_state(&format!("{p}.ffn.value.weight"), &mut b.ffn.value_w)?;
1834 }
1835 Ok(adam)
1836 }
1837
1838 pub fn config(&self) -> &Config {
1840 &self.cfg
1841 }
1842
1843 pub fn new_state(&self) -> State {
1845 State::new(&self.cfg)
1846 }
1847
1848 #[inline]
1850 pub fn lm_head_weights(&self) -> &[f32] {
1851 self.lm_head.as_slice()
1852 }
1853
1854 #[inline]
1856 pub fn lm_head_weights_mut(&mut self) -> &mut [f32] {
1857 self.lm_head.as_mut_slice()
1858 }
1859
1860 #[allow(clippy::too_many_arguments)]
1861 fn apply_full_gradients(
1862 &mut self,
1863 grads: &FullGradState,
1864 scope: TrainScopeMask,
1865 optimizer: OptimizerKind,
1866 lr: f32,
1867 clip: f32,
1868 adam_t: &mut usize,
1869 model_adam: Option<&mut FullAdamState>,
1870 out_bias: Option<&mut [f32]>,
1871 out_bias_grad: Option<&[f32]>,
1872 out_bias_adam_m: Option<&mut [f32]>,
1873 out_bias_adam_v: Option<&mut [f32]>,
1874 ) -> Result<()> {
1875 let mut adam_step = None::<AdamStep>;
1876 let mut model_adam = model_adam;
1877 if matches!(optimizer, OptimizerKind::Adam) {
1878 *adam_t = adam_t.saturating_add(1);
1879 let t = (*adam_t).max(1) as i32;
1880 let b1 = 0.9f32;
1881 let b2 = 0.999f32;
1882 adam_step = Some(AdamStep {
1883 lr,
1884 clip: clip.max(0.0),
1885 b1,
1886 b2,
1887 eps: 1e-8,
1888 bias_corr1: 1.0 - b1.powi(t),
1889 bias_corr2: 1.0 - b2.powi(t),
1890 });
1891 if scope.trains_non_head_params() && model_adam.is_none() {
1892 bail!("rwkv Adam full-training state is missing");
1893 }
1894 }
1895
1896 if scope.bias
1897 && let (Some(bias), Some(grad)) = (out_bias, out_bias_grad)
1898 {
1899 match optimizer {
1900 OptimizerKind::Sgd => sgd_vec_update(bias, grad, lr, clip),
1901 OptimizerKind::Adam => {
1902 let cfg = adam_step.as_ref().expect("adam cfg initialized");
1903 let Some(m) = out_bias_adam_m else {
1904 bail!("rwkv Adam output-bias state is missing (m)");
1905 };
1906 let Some(v) = out_bias_adam_v else {
1907 bail!("rwkv Adam output-bias state is missing (v)");
1908 };
1909 apply_adam_vec_update_raw(bias, grad, m, v, cfg);
1910 }
1911 }
1912 }
1913
1914 if scope.head {
1915 match optimizer {
1916 OptimizerKind::Sgd => {
1917 sgd_vec_update(
1918 self.lm_head.as_mut_slice(),
1919 grads.lm_head.as_slice(),
1920 lr,
1921 clip,
1922 );
1923 sgd_vec_update(
1924 self.ln_out_w.as_mut_slice(),
1925 grads.ln_out_w.as_slice(),
1926 lr,
1927 clip,
1928 );
1929 sgd_vec_update(
1930 self.ln_out_b.as_mut_slice(),
1931 grads.ln_out_b.as_slice(),
1932 lr,
1933 clip,
1934 );
1935 }
1936 OptimizerKind::Adam => {
1937 let cfg = adam_step.as_ref().expect("adam cfg initialized");
1938 let adam = model_adam.as_mut().expect("adam state exists");
1939 apply_adam_vec_update(
1940 self.lm_head.as_mut_slice(),
1941 grads.lm_head.as_slice(),
1942 &mut adam.lm_head,
1943 cfg,
1944 );
1945 apply_adam_vec_update(
1946 self.ln_out_w.as_mut_slice(),
1947 grads.ln_out_w.as_slice(),
1948 &mut adam.ln_out_w,
1949 cfg,
1950 );
1951 apply_adam_vec_update(
1952 self.ln_out_b.as_mut_slice(),
1953 grads.ln_out_b.as_slice(),
1954 &mut adam.ln_out_b,
1955 cfg,
1956 );
1957 }
1958 }
1959 }
1960
1961 for layer_idx in 0..self.cfg.num_layers {
1962 let block = &mut self.blocks[layer_idx];
1963 let grad = &grads.blocks[layer_idx];
1964 match optimizer {
1965 OptimizerKind::Sgd => {
1966 if scope.ffn {
1967 sgd_vec_update(
1968 block.ffn.x_k.as_mut_slice(),
1969 grad.ffn.x_k.as_slice(),
1970 lr,
1971 clip,
1972 );
1973 sgd_vec_update(
1974 block.ffn.key_w.as_mut_slice(),
1975 grad.ffn.key_w.as_slice(),
1976 lr,
1977 clip,
1978 );
1979 sgd_vec_update(
1980 block.ffn.value_w.as_mut_slice(),
1981 grad.ffn.value_w.as_slice(),
1982 lr,
1983 clip,
1984 );
1985 }
1986 if scope.ffn_norm {
1987 sgd_vec_update(
1988 block.ffn_norm_w.as_mut_slice(),
1989 grad.ffn_norm_w.as_slice(),
1990 lr,
1991 clip,
1992 );
1993 sgd_vec_update(
1994 block.ffn_norm_b.as_mut_slice(),
1995 grad.ffn_norm_b.as_slice(),
1996 lr,
1997 clip,
1998 );
1999 }
2000 if scope.attn {
2001 sgd_vec_update(
2002 block.attn.o_proj.as_mut_slice(),
2003 grad.attn.o_proj.as_slice(),
2004 lr,
2005 clip,
2006 );
2007 sgd_vec_update(
2008 block.attn.r_k.as_mut_slice(),
2009 grad.attn.r_k.as_slice(),
2010 lr,
2011 clip,
2012 );
2013 sgd_vec_update(
2014 block.attn.g_norm_w.as_mut_slice(),
2015 grad.attn.g_norm_w.as_slice(),
2016 lr,
2017 clip,
2018 );
2019 sgd_vec_update(
2020 block.attn.g_norm_b.as_mut_slice(),
2021 grad.attn.g_norm_b.as_slice(),
2022 lr,
2023 clip,
2024 );
2025 sgd_vec_update(
2026 block.attn.k_a.as_mut_slice(),
2027 grad.attn.k_a.as_slice(),
2028 lr,
2029 clip,
2030 );
2031 sgd_vec_update(
2032 block.attn.k_k.as_mut_slice(),
2033 grad.attn.k_k.as_slice(),
2034 lr,
2035 clip,
2036 );
2037 sgd_vec_update(
2038 block.attn.rkv_proj.as_mut_slice(),
2039 grad.attn.rkv_proj.as_slice(),
2040 lr,
2041 clip,
2042 );
2043 sgd_vec_update(
2044 block.attn.w0.as_mut_slice(),
2045 grad.attn.w0.as_slice(),
2046 lr,
2047 clip,
2048 );
2049 sgd_vec_update(
2050 block.attn.w2.as_mut_slice(),
2051 grad.attn.w2.as_slice(),
2052 lr,
2053 clip,
2054 );
2055 sgd_vec_update(
2056 block.attn.w1.as_mut_slice(),
2057 grad.attn.w1.as_slice(),
2058 lr,
2059 clip,
2060 );
2061 sgd_vec_update(
2062 block.attn.a0.as_mut_slice(),
2063 grad.attn.a0.as_slice(),
2064 lr,
2065 clip,
2066 );
2067 sgd_vec_update(
2068 block.attn.a2.as_mut_slice(),
2069 grad.attn.a2.as_slice(),
2070 lr,
2071 clip,
2072 );
2073 sgd_vec_update(
2074 block.attn.a1.as_mut_slice(),
2075 grad.attn.a1.as_slice(),
2076 lr,
2077 clip,
2078 );
2079 sgd_vec_update(
2080 block.attn.g2.as_mut_slice(),
2081 grad.attn.g2.as_slice(),
2082 lr,
2083 clip,
2084 );
2085 sgd_vec_update(
2086 block.attn.g1.as_mut_slice(),
2087 grad.attn.g1.as_slice(),
2088 lr,
2089 clip,
2090 );
2091 sgd_vec_update(
2092 block.attn.x_r.as_mut_slice(),
2093 grad.attn.x_r.as_slice(),
2094 lr,
2095 clip,
2096 );
2097 sgd_vec_update(
2098 block.attn.x_w.as_mut_slice(),
2099 grad.attn.x_w.as_slice(),
2100 lr,
2101 clip,
2102 );
2103 sgd_vec_update(
2104 block.attn.x_k.as_mut_slice(),
2105 grad.attn.x_k.as_slice(),
2106 lr,
2107 clip,
2108 );
2109 sgd_vec_update(
2110 block.attn.x_v.as_mut_slice(),
2111 grad.attn.x_v.as_slice(),
2112 lr,
2113 clip,
2114 );
2115 sgd_vec_update(
2116 block.attn.x_a.as_mut_slice(),
2117 grad.attn.x_a.as_slice(),
2118 lr,
2119 clip,
2120 );
2121 sgd_vec_update(
2122 block.attn.x_g.as_mut_slice(),
2123 grad.attn.x_g.as_slice(),
2124 lr,
2125 clip,
2126 );
2127 if let (Some(v1), Some(gv1)) =
2128 (block.attn.v1.as_mut(), grad.attn.v1.as_ref())
2129 {
2130 sgd_vec_update(v1.as_mut_slice(), gv1.as_slice(), lr, clip);
2131 }
2132 if let (Some(v2), Some(gv2)) =
2133 (block.attn.v2.as_mut(), grad.attn.v2.as_ref())
2134 {
2135 sgd_vec_update(v2.as_mut_slice(), gv2.as_slice(), lr, clip);
2136 }
2137 if let (Some(v0), Some(gv0)) =
2138 (block.attn.v0.as_mut(), grad.attn.v0.as_ref())
2139 {
2140 sgd_vec_update(v0.as_mut_slice(), gv0.as_slice(), lr, clip);
2141 }
2142 }
2143 if scope.attn_norm {
2144 sgd_vec_update(
2145 block.attn_norm_w.as_mut_slice(),
2146 grad.attn_norm_w.as_slice(),
2147 lr,
2148 clip,
2149 );
2150 sgd_vec_update(
2151 block.attn_norm_b.as_mut_slice(),
2152 grad.attn_norm_b.as_slice(),
2153 lr,
2154 clip,
2155 );
2156 }
2157 if scope.pre_norm
2158 && let (Some(w), Some(gw)) =
2159 (block.pre_norm_w.as_mut(), grad.pre_norm_w.as_ref())
2160 {
2161 sgd_vec_update(w.as_mut_slice(), gw.as_slice(), lr, clip);
2162 }
2163 if scope.pre_norm
2164 && let (Some(b), Some(gb)) =
2165 (block.pre_norm_b.as_mut(), grad.pre_norm_b.as_ref())
2166 {
2167 sgd_vec_update(b.as_mut_slice(), gb.as_slice(), lr, clip);
2168 }
2169 }
2170 OptimizerKind::Adam => {
2171 let cfg = adam_step.as_ref().expect("adam cfg initialized");
2172 let adam =
2173 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
2174 if scope.ffn {
2175 apply_adam_vec_update(
2176 block.ffn.x_k.as_mut_slice(),
2177 grad.ffn.x_k.as_slice(),
2178 &mut adam.ffn.x_k,
2179 cfg,
2180 );
2181 apply_adam_vec_update(
2182 block.ffn.key_w.as_mut_slice(),
2183 grad.ffn.key_w.as_slice(),
2184 &mut adam.ffn.key_w,
2185 cfg,
2186 );
2187 apply_adam_vec_update(
2188 block.ffn.value_w.as_mut_slice(),
2189 grad.ffn.value_w.as_slice(),
2190 &mut adam.ffn.value_w,
2191 cfg,
2192 );
2193 }
2194 if scope.ffn_norm {
2195 apply_adam_vec_update(
2196 block.ffn_norm_w.as_mut_slice(),
2197 grad.ffn_norm_w.as_slice(),
2198 &mut adam.ffn_norm_w,
2199 cfg,
2200 );
2201 apply_adam_vec_update(
2202 block.ffn_norm_b.as_mut_slice(),
2203 grad.ffn_norm_b.as_slice(),
2204 &mut adam.ffn_norm_b,
2205 cfg,
2206 );
2207 }
2208 if scope.attn {
2209 apply_adam_vec_update(
2210 block.attn.o_proj.as_mut_slice(),
2211 grad.attn.o_proj.as_slice(),
2212 &mut adam.attn.o_proj,
2213 cfg,
2214 );
2215 apply_adam_vec_update(
2216 block.attn.r_k.as_mut_slice(),
2217 grad.attn.r_k.as_slice(),
2218 &mut adam.attn.r_k,
2219 cfg,
2220 );
2221 apply_adam_vec_update(
2222 block.attn.g_norm_w.as_mut_slice(),
2223 grad.attn.g_norm_w.as_slice(),
2224 &mut adam.attn.g_norm_w,
2225 cfg,
2226 );
2227 apply_adam_vec_update(
2228 block.attn.g_norm_b.as_mut_slice(),
2229 grad.attn.g_norm_b.as_slice(),
2230 &mut adam.attn.g_norm_b,
2231 cfg,
2232 );
2233 apply_adam_vec_update(
2234 block.attn.k_a.as_mut_slice(),
2235 grad.attn.k_a.as_slice(),
2236 &mut adam.attn.k_a,
2237 cfg,
2238 );
2239 apply_adam_vec_update(
2240 block.attn.k_k.as_mut_slice(),
2241 grad.attn.k_k.as_slice(),
2242 &mut adam.attn.k_k,
2243 cfg,
2244 );
2245 apply_adam_vec_update(
2246 block.attn.rkv_proj.as_mut_slice(),
2247 grad.attn.rkv_proj.as_slice(),
2248 &mut adam.attn.rkv_proj,
2249 cfg,
2250 );
2251 apply_adam_vec_update(
2252 block.attn.w0.as_mut_slice(),
2253 grad.attn.w0.as_slice(),
2254 &mut adam.attn.w0,
2255 cfg,
2256 );
2257 apply_adam_vec_update(
2258 block.attn.w2.as_mut_slice(),
2259 grad.attn.w2.as_slice(),
2260 &mut adam.attn.w2,
2261 cfg,
2262 );
2263 apply_adam_vec_update(
2264 block.attn.w1.as_mut_slice(),
2265 grad.attn.w1.as_slice(),
2266 &mut adam.attn.w1,
2267 cfg,
2268 );
2269 apply_adam_vec_update(
2270 block.attn.a0.as_mut_slice(),
2271 grad.attn.a0.as_slice(),
2272 &mut adam.attn.a0,
2273 cfg,
2274 );
2275 apply_adam_vec_update(
2276 block.attn.a2.as_mut_slice(),
2277 grad.attn.a2.as_slice(),
2278 &mut adam.attn.a2,
2279 cfg,
2280 );
2281 apply_adam_vec_update(
2282 block.attn.a1.as_mut_slice(),
2283 grad.attn.a1.as_slice(),
2284 &mut adam.attn.a1,
2285 cfg,
2286 );
2287 apply_adam_vec_update(
2288 block.attn.g2.as_mut_slice(),
2289 grad.attn.g2.as_slice(),
2290 &mut adam.attn.g2,
2291 cfg,
2292 );
2293 apply_adam_vec_update(
2294 block.attn.g1.as_mut_slice(),
2295 grad.attn.g1.as_slice(),
2296 &mut adam.attn.g1,
2297 cfg,
2298 );
2299 apply_adam_vec_update(
2300 block.attn.x_r.as_mut_slice(),
2301 grad.attn.x_r.as_slice(),
2302 &mut adam.attn.x_r,
2303 cfg,
2304 );
2305 apply_adam_vec_update(
2306 block.attn.x_w.as_mut_slice(),
2307 grad.attn.x_w.as_slice(),
2308 &mut adam.attn.x_w,
2309 cfg,
2310 );
2311 apply_adam_vec_update(
2312 block.attn.x_k.as_mut_slice(),
2313 grad.attn.x_k.as_slice(),
2314 &mut adam.attn.x_k,
2315 cfg,
2316 );
2317 apply_adam_vec_update(
2318 block.attn.x_v.as_mut_slice(),
2319 grad.attn.x_v.as_slice(),
2320 &mut adam.attn.x_v,
2321 cfg,
2322 );
2323 apply_adam_vec_update(
2324 block.attn.x_a.as_mut_slice(),
2325 grad.attn.x_a.as_slice(),
2326 &mut adam.attn.x_a,
2327 cfg,
2328 );
2329 apply_adam_vec_update(
2330 block.attn.x_g.as_mut_slice(),
2331 grad.attn.x_g.as_slice(),
2332 &mut adam.attn.x_g,
2333 cfg,
2334 );
2335 if let (Some(v1), Some(gv1), Some(av1)) = (
2336 block.attn.v1.as_mut(),
2337 grad.attn.v1.as_ref(),
2338 adam.attn.v1.as_mut(),
2339 ) {
2340 apply_adam_vec_update(v1.as_mut_slice(), gv1.as_slice(), av1, cfg);
2341 }
2342 if let (Some(v2), Some(gv2), Some(av2)) = (
2343 block.attn.v2.as_mut(),
2344 grad.attn.v2.as_ref(),
2345 adam.attn.v2.as_mut(),
2346 ) {
2347 apply_adam_vec_update(v2.as_mut_slice(), gv2.as_slice(), av2, cfg);
2348 }
2349 if let (Some(v0), Some(gv0), Some(av0)) = (
2350 block.attn.v0.as_mut(),
2351 grad.attn.v0.as_ref(),
2352 adam.attn.v0.as_mut(),
2353 ) {
2354 apply_adam_vec_update(v0.as_mut_slice(), gv0.as_slice(), av0, cfg);
2355 }
2356 }
2357 if scope.attn_norm {
2358 apply_adam_vec_update(
2359 block.attn_norm_w.as_mut_slice(),
2360 grad.attn_norm_w.as_slice(),
2361 &mut adam.attn_norm_w,
2362 cfg,
2363 );
2364 apply_adam_vec_update(
2365 block.attn_norm_b.as_mut_slice(),
2366 grad.attn_norm_b.as_slice(),
2367 &mut adam.attn_norm_b,
2368 cfg,
2369 );
2370 }
2371 if scope.pre_norm
2372 && let (Some(w), Some(gw), Some(aw)) = (
2373 block.pre_norm_w.as_mut(),
2374 grad.pre_norm_w.as_ref(),
2375 adam.pre_norm_w.as_mut(),
2376 )
2377 {
2378 apply_adam_vec_update(w.as_mut_slice(), gw.as_slice(), aw, cfg);
2379 }
2380 if scope.pre_norm
2381 && let (Some(b), Some(gb), Some(ab)) = (
2382 block.pre_norm_b.as_mut(),
2383 grad.pre_norm_b.as_ref(),
2384 adam.pre_norm_b.as_mut(),
2385 )
2386 {
2387 apply_adam_vec_update(b.as_mut_slice(), gb.as_slice(), ab, cfg);
2388 }
2389 }
2390 }
2391 }
2392
2393 if scope.embed {
2394 match optimizer {
2395 OptimizerKind::Sgd => {
2396 sgd_vec_update(
2397 self.embeddings.as_mut_slice(),
2398 grads.embeddings.as_slice(),
2399 lr,
2400 clip,
2401 );
2402 }
2403 OptimizerKind::Adam => {
2404 let cfg = adam_step.as_ref().expect("adam cfg initialized");
2405 let adam = model_adam.as_mut().expect("adam state exists");
2406 apply_adam_vec_update(
2407 self.embeddings.as_mut_slice(),
2408 grads.embeddings.as_slice(),
2409 &mut adam.embeddings,
2410 cfg,
2411 );
2412 }
2413 }
2414 }
2415 Ok(())
2416 }
2417
2418 #[allow(clippy::needless_range_loop, clippy::too_many_arguments)]
2419 fn accumulate_token_step_gradients(
2420 &self,
2421 scratch: &mut ScratchBuffers,
2422 trace: &TokenTrainTrace,
2423 state_new: &State,
2424 symbol: u8,
2425 pdf: &[f64],
2426 grad_scale: f32,
2427 scope: TrainScopeMask,
2428 grads: &mut FullGradState,
2429 out_bias_grad: Option<&mut [f32]>,
2430 future: &mut RecurrentGradState,
2431 ) -> Result<()> {
2432 let c = self.cfg.hidden_size;
2433 let h = self.cfg.num_heads;
2434 let n = self.cfg.head_dim;
2435 let i = self.cfg.intermediate_size;
2436 let d_w = self.cfg.decay_low_rank;
2437 let d_a = self.cfg.a_low_rank;
2438 let d_v = self.cfg.v_low_rank;
2439 let d_g = self.cfg.g_low_rank;
2440 let vocab = self.cfg.vocab_size.min(pdf.len());
2441 if vocab == 0 {
2442 return Ok(());
2443 }
2444
2445 scratch.grad_logits.zero();
2446 for idx in 0..vocab {
2447 let p = pdf[idx].clamp(1e-12, 1.0) as f32;
2448 let target = if idx == symbol as usize { 1.0 } else { 0.0 };
2449 scratch.grad_logits[idx] = (target - p) * grad_scale;
2450 }
2451
2452 if scope.bias
2453 && let Some(bias_grad) = out_bias_grad
2454 {
2455 add_vec_grad(
2456 &mut bias_grad[0..vocab],
2457 &scratch.grad_logits.as_slice()[0..vocab],
2458 );
2459 }
2460
2461 scratch.grad_x.zero();
2462 if scope.head {
2463 add_outer_grad(
2464 grads.lm_head.as_mut_slice(),
2465 vocab,
2466 c,
2467 &scratch.grad_logits.as_slice()[0..vocab],
2468 trace.x_normed.as_slice(),
2469 );
2470 }
2471 for row in 0..vocab {
2472 let g = scratch.grad_logits[row];
2473 if g == 0.0 {
2474 continue;
2475 }
2476 let row_off = row * c;
2477 for col in 0..c {
2478 scratch.grad_x[col] += self.lm_head[row_off + col] * g;
2479 }
2480 }
2481
2482 let needs_backprop = scope.trains_non_head_params() || scope.head;
2483 if !needs_backprop {
2484 return Ok(());
2485 }
2486
2487 layer_norm_backward(
2488 trace.x.as_slice(),
2489 self.ln_out_w.as_slice(),
2490 scratch.grad_x.as_slice(),
2491 self.cfg.layer_norm_eps,
2492 scratch.grad_x2.as_mut_slice(),
2493 scratch.grad_x3.as_mut_slice(),
2494 scratch.grad_x4.as_mut_slice(),
2495 );
2496 if scope.head {
2497 add_vec_grad(grads.ln_out_w.as_mut_slice(), scratch.grad_x3.as_slice());
2498 add_vec_grad(grads.ln_out_b.as_mut_slice(), scratch.grad_x4.as_slice());
2499 }
2500 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice());
2501 scratch.grad_v_first.zero();
2502
2503 for layer_idx in (0..self.cfg.num_layers).rev() {
2504 let tr = &trace.layers[layer_idx];
2505 let block = &self.blocks[layer_idx];
2506 let block_grads = &mut grads.blocks[layer_idx];
2507 let future_layer = &mut future.layers[layer_idx];
2508
2509 scratch.grad_x2.copy_from_slice(scratch.grad_x.as_slice());
2510 scratch.grad_x3.copy_from_slice(scratch.grad_x.as_slice());
2511
2512 unsafe {
2513 kernel::gemv_t_avx(
2514 block.ffn.value_w.as_ptr(),
2515 scratch.grad_x3.as_ptr(),
2516 scratch.grad_ffn.as_mut_ptr(),
2517 c,
2518 i,
2519 );
2520 }
2521 if scope.ffn {
2522 add_outer_grad(
2523 block_grads.ffn.value_w.as_mut_slice(),
2524 c,
2525 i,
2526 scratch.grad_x3.as_slice(),
2527 tr.ffn_k.as_slice(),
2528 );
2529 }
2530
2531 for col in 0..i {
2532 let pre = tr.ffn_pre[col];
2533 scratch.grad_ffn2[col] = if pre > 0.0 {
2534 scratch.grad_ffn[col] * (2.0 * pre)
2535 } else {
2536 0.0
2537 };
2538 }
2539
2540 unsafe {
2541 kernel::gemv_t_avx(
2542 block.ffn.key_w.as_ptr(),
2543 scratch.grad_ffn2.as_ptr(),
2544 scratch.grad_x4.as_mut_ptr(),
2545 i,
2546 c,
2547 );
2548 }
2549 if scope.ffn {
2550 add_outer_grad(
2551 block_grads.ffn.key_w.as_mut_slice(),
2552 i,
2553 c,
2554 scratch.grad_ffn2.as_slice(),
2555 tr.ffn_xk.as_slice(),
2556 );
2557 }
2558
2559 scratch
2560 .grad_x5
2561 .copy_from_slice(future_layer.ffn_x_prev.as_slice());
2562 future_layer.ffn_x_prev.zero();
2563 for col in 0..c {
2564 let g = scratch.grad_x4[col];
2565 let mix = block.ffn.x_k[col];
2566 let base = tr.ffn_norm[col];
2567 let prev = tr.ffn_x_prev_old[col];
2568 scratch.grad_x5[col] += g * (1.0 - mix);
2569 future_layer.ffn_x_prev[col] = g * mix;
2570 scratch.grad_param[col] = g * (prev - base);
2571 }
2572 if scope.ffn {
2573 add_vec_grad(
2574 block_grads.ffn.x_k.as_mut_slice(),
2575 scratch.grad_param.as_slice(),
2576 );
2577 }
2578
2579 layer_norm_backward(
2580 tr.x_after_attn.as_slice(),
2581 block.ffn_norm_w.as_slice(),
2582 scratch.grad_x5.as_slice(),
2583 self.cfg.layer_norm_eps,
2584 scratch.grad_x4.as_mut_slice(),
2585 scratch.grad_x3.as_mut_slice(),
2586 scratch.grad_x6.as_mut_slice(),
2587 );
2588 if scope.ffn_norm {
2589 add_vec_grad(
2590 block_grads.ffn_norm_w.as_mut_slice(),
2591 scratch.grad_x3.as_slice(),
2592 );
2593 add_vec_grad(
2594 block_grads.ffn_norm_b.as_mut_slice(),
2595 scratch.grad_x6.as_slice(),
2596 );
2597 }
2598 for col in 0..c {
2599 scratch.grad_x2[col] += scratch.grad_x4[col];
2600 }
2601
2602 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice());
2603 scratch.grad_x3.copy_from_slice(scratch.grad_x2.as_slice());
2604
2605 unsafe {
2606 kernel::gemv_t_avx(
2607 block.attn.o_proj.as_ptr(),
2608 scratch.grad_x3.as_ptr(),
2609 scratch.grad_x4.as_mut_ptr(),
2610 c,
2611 c,
2612 );
2613 }
2614 if scope.attn {
2615 add_outer_grad(
2616 block_grads.attn.o_proj.as_mut_slice(),
2617 c,
2618 c,
2619 scratch.grad_x3.as_slice(),
2620 tr.y_gate.as_slice(),
2621 );
2622 }
2623
2624 for col in 0..c {
2625 let gy = scratch.grad_x4[col];
2626 scratch.grad_saved[col] = gy * tr.y_head[col];
2627 scratch.grad_x4[col] = gy * tr.g[col];
2628 }
2629
2630 scratch.grad_x2.zero();
2631 scratch.grad_x3.zero();
2632 scratch.grad_x6.zero();
2633 scratch.grad_param.zero();
2634 for head_idx in 0..h {
2635 let off = head_idx * n;
2636 let mut g_alpha = 0.0f32;
2637 for j in 0..n {
2638 let g = scratch.grad_x4[off + j];
2639 g_alpha += g * tr.v[off + j];
2640 scratch.grad_x6[off + j] += g * tr.alpha[head_idx];
2641 }
2642 for j in 0..n {
2643 let idx = off + j;
2644 let rk = block.attn.r_k[idx];
2645 let rv = tr.r[idx];
2646 let kv = tr.k[idx];
2647 let g = g_alpha * rk;
2648 scratch.grad_x2[idx] += g * kv;
2649 scratch.grad_x3[idx] += g * rv;
2650 scratch.grad_param[idx] += g_alpha * rv * kv;
2651 }
2652 }
2653 if scope.attn {
2654 add_vec_grad(
2655 block_grads.attn.r_k.as_mut_slice(),
2656 scratch.grad_param.as_slice(),
2657 );
2658 }
2659
2660 scratch.grad_x5.as_mut_slice()[0..c].copy_from_slice(&scratch.grad_x4.as_slice()[0..c]);
2661 group_norm_backward(
2662 tr.y_wkv.as_slice(),
2663 block.attn.g_norm_w.as_slice(),
2664 scratch.grad_x5.as_slice(),
2665 h,
2666 n,
2667 self.cfg.group_norm_eps,
2668 scratch.grad_x4.as_mut_slice(),
2669 scratch.grad_param.as_mut_slice(),
2670 scratch.grad_param2.as_mut_slice(),
2671 );
2672 if scope.attn {
2673 add_vec_grad(
2674 block_grads.attn.g_norm_w.as_mut_slice(),
2675 scratch.grad_param.as_slice(),
2676 );
2677 add_vec_grad(
2678 block_grads.attn.g_norm_b.as_mut_slice(),
2679 scratch.grad_param2.as_slice(),
2680 );
2681 }
2682
2683 scratch.grad_param.zero();
2684 scratch.grad_x5.zero();
2685 scratch.grad_param2.zero();
2686 scratch
2687 .grad_att_state
2688 .copy_from_slice(future_layer.att_state.as_slice());
2689 future_layer.att_state.zero();
2690 let s_old = tr.att_state_old.as_slice();
2691 let s_new = state_new.layers[layer_idx].att_state.as_slice();
2692 for head_idx in 0..h {
2693 let off = head_idx * n;
2694 let s_off = head_idx * n * n;
2695 let grad_y = &scratch.grad_x4.as_slice()[off..off + n];
2696 let r_head = &tr.r.as_slice()[off..off + n];
2697 let k_head = &tr.k.as_slice()[off..off + n];
2698 let kk_head = &tr.kk.as_slice()[off..off + n];
2699 let a_head = &tr.a.as_slice()[off..off + n];
2700 let v_head = &tr.v.as_slice()[off..off + n];
2701 let w_head = &tr.w_decay.as_slice()[off..off + n];
2702
2703 unsafe {
2704 kernel::gemv_t_avx(
2705 s_new.as_ptr().add(s_off),
2706 grad_y.as_ptr(),
2707 scratch.grad_low_rank.as_mut_ptr(),
2708 n,
2709 n,
2710 );
2711 }
2712 for j in 0..n {
2713 scratch.grad_x2[off + j] += scratch.grad_low_rank[j];
2714 }
2715
2716 let g_state = &mut scratch.grad_att_state.as_mut_slice()[s_off..s_off + n * n];
2717 for irow in 0..n {
2718 let gy = grad_y[irow];
2719 let row_off = irow * n;
2720 for j in 0..n {
2721 g_state[row_off + j] += gy * r_head[j];
2722 }
2723 }
2724
2725 unsafe {
2726 kernel::gemv_avx(
2727 s_old.as_ptr().add(s_off),
2728 kk_head.as_ptr(),
2729 scratch.grad_low_rank.as_mut_ptr(),
2730 n,
2731 n,
2732 );
2733 }
2734 let u = &scratch.grad_low_rank.as_slice()[0..n];
2735
2736 for j in 0..n {
2737 let mut grad_w = 0.0f32;
2738 let mut grad_k = 0.0f32;
2739 let mut grad_b = 0.0f32;
2740 for irow in 0..n {
2741 let g = g_state[irow * n + j];
2742 grad_w += g * s_old[s_off + irow * n + j];
2743 grad_k += g * v_head[irow];
2744 grad_b -= g * u[irow];
2745 future_layer.att_state[s_off + irow * n + j] = g * w_head[j];
2746 }
2747 scratch.grad_param[off + j] += grad_w;
2748 scratch.grad_x3[off + j] += grad_k;
2749 scratch.grad_param2[off + j] += grad_b * a_head[j];
2750 scratch.grad_x5[off + j] += grad_b * kk_head[j];
2751 }
2752
2753 for irow in 0..n {
2754 let mut grad_u = 0.0f32;
2755 for j in 0..n {
2756 grad_u -= g_state[irow * n + j] * kk_head[j] * a_head[j];
2757 }
2758 scratch.grad_low_rank2[irow] = grad_u;
2759 let row_off = irow * n;
2760 for j in 0..n {
2761 future_layer.att_state[s_off + row_off + j] += grad_u * kk_head[j];
2762 }
2763 }
2764 unsafe {
2765 kernel::gemv_t_avx(
2766 s_old.as_ptr().add(s_off),
2767 scratch.grad_low_rank2.as_ptr(),
2768 scratch.grad_low_rank.as_mut_ptr(),
2769 n,
2770 n,
2771 );
2772 }
2773 for j in 0..n {
2774 scratch.grad_param2[off + j] += scratch.grad_low_rank[j];
2775 }
2776
2777 for irow in 0..n {
2778 let mut grad_v = 0.0f32;
2779 for j in 0..n {
2780 grad_v += g_state[irow * n + j] * k_head[j];
2781 }
2782 scratch.grad_x6[off + irow] += grad_v;
2783 }
2784 }
2785
2786 for col in 0..c {
2787 let gk = scratch.grad_x3[col];
2788 let scale = 1.0 + (tr.a[col] - 1.0) * block.attn.k_a[col];
2789 let d_scale = gk * tr.k_pre[col];
2790 scratch.grad_x3[col] = gk * scale;
2791 scratch.grad_x5[col] += d_scale * block.attn.k_a[col];
2792 scratch.grad_param[col] = d_scale * (tr.a[col] - 1.0);
2793 }
2794 for head_idx in 0..h {
2795 let off = head_idx * n;
2796 l2_normalize_backward(
2797 &tr.kk_pre.as_slice()[off..off + n],
2798 &tr.kk.as_slice()[off..off + n],
2799 &scratch.grad_param2.as_slice()[off..off + n],
2800 1e-12,
2801 &mut scratch.grad_x4.as_mut_slice()[off..off + n],
2802 );
2803 }
2804 for col in 0..c {
2805 let g = scratch.grad_x4[col];
2806 scratch.grad_x3[col] += g * block.attn.k_k[col];
2807 scratch.grad_param2[col] = g * tr.k_pre[col];
2808 }
2809 if scope.attn {
2810 add_vec_grad(
2811 block_grads.attn.k_a.as_mut_slice(),
2812 scratch.grad_param.as_slice(),
2813 );
2814 add_vec_grad(
2815 block_grads.attn.k_k.as_mut_slice(),
2816 scratch.grad_param2.as_slice(),
2817 );
2818 }
2819
2820 scratch
2821 .grad_param2
2822 .copy_from_slice(scratch.grad_x6.as_slice());
2823 if layer_idx == 0 {
2824 for col in 0..c {
2825 scratch.grad_x6[col] += scratch.grad_v_first[col];
2826 }
2827 } else if tr.uses_v_residual
2828 && block.attn.v1.is_some()
2829 && block.attn.v2.is_some()
2830 && block.attn.v0.is_some()
2831 {
2832 let v1 = block.attn.v1.as_ref().expect("v1");
2833 let v2 = block.attn.v2.as_ref().expect("v2");
2834 for col in 0..c {
2835 let gv = scratch.grad_param2[col];
2836 let nu = tr.nu[col];
2837 scratch.grad_x6[col] = gv * (1.0 - nu);
2838 scratch.grad_x3[col] = gv * (trace.v_first[col] - tr.v_pre[col]);
2839 scratch.grad_v_first[col] += gv * nu;
2840 }
2841 for col in 0..c {
2842 let nu = tr.nu[col];
2843 scratch.grad_x3[col] *= nu * (1.0 - nu);
2844 }
2845 if scope.attn {
2846 add_vec_grad(
2847 block_grads
2848 .attn
2849 .v0
2850 .as_mut()
2851 .expect("grad v0")
2852 .as_mut_slice(),
2853 scratch.grad_x3.as_slice(),
2854 );
2855 add_outer_grad(
2856 block_grads
2857 .attn
2858 .v2
2859 .as_mut()
2860 .expect("grad v2")
2861 .as_mut_slice(),
2862 c,
2863 d_v,
2864 scratch.grad_x3.as_slice(),
2865 &tr.v_hidden.as_slice()[0..d_v],
2866 );
2867 }
2868 unsafe {
2869 kernel::gemv_t_avx(
2870 v2.as_ptr(),
2871 scratch.grad_x3.as_ptr(),
2872 scratch.grad_low_rank.as_mut_ptr(),
2873 c,
2874 d_v,
2875 );
2876 }
2877 if scope.attn {
2878 add_outer_grad(
2879 block_grads
2880 .attn
2881 .v1
2882 .as_mut()
2883 .expect("grad v1")
2884 .as_mut_slice(),
2885 d_v,
2886 c,
2887 &scratch.grad_low_rank.as_slice()[0..d_v],
2888 tr.xv.as_slice(),
2889 );
2890 }
2891 for col in 0..c {
2892 let mut acc = 0.0f32;
2893 for row in 0..d_v {
2894 acc += v1[row * c + col] * scratch.grad_low_rank[row];
2895 }
2896 scratch.grad_x4[col] += acc;
2897 }
2898 }
2899
2900 let proj_size = c * c;
2901 if scope.attn {
2902 add_outer_grad(
2903 &mut block_grads.attn.rkv_proj.as_mut_slice()[0..proj_size],
2904 c,
2905 c,
2906 scratch.grad_x2.as_slice(),
2907 tr.xr.as_slice(),
2908 );
2909 add_outer_grad(
2910 &mut block_grads.attn.rkv_proj.as_mut_slice()[proj_size..2 * proj_size],
2911 c,
2912 c,
2913 scratch.grad_x3.as_slice(),
2914 tr.xk.as_slice(),
2915 );
2916 add_outer_grad(
2917 &mut block_grads.attn.rkv_proj.as_mut_slice()[2 * proj_size..3 * proj_size],
2918 c,
2919 c,
2920 scratch.grad_x6.as_slice(),
2921 tr.xv.as_slice(),
2922 );
2923 }
2924 let proj = block.attn.rkv_proj.as_slice();
2925 unsafe {
2926 kernel::gemv_t_avx(
2927 proj.as_ptr(),
2928 scratch.grad_x2.as_ptr(),
2929 scratch.grad_param.as_mut_ptr(),
2930 c,
2931 c,
2932 );
2933 kernel::gemv_t_avx(
2934 proj.as_ptr().add(proj_size),
2935 scratch.grad_x3.as_ptr(),
2936 scratch.grad_param2.as_mut_ptr(),
2937 c,
2938 c,
2939 );
2940 kernel::gemv_t_avx(
2941 proj.as_ptr().add(2 * proj_size),
2942 scratch.grad_x6.as_ptr(),
2943 scratch.grad_x4.as_mut_ptr(),
2944 c,
2945 c,
2946 );
2947 }
2948
2949 let inv_sqrt_e = 1.0 / std::f32::consts::E.sqrt();
2950 for col in 0..c {
2951 let sig = tr.w_sigmoid[col];
2952 let d_sig = scratch.grad_param[col] * (-inv_sqrt_e) * tr.w_decay[col];
2953 scratch.grad_param[col] = d_sig * sig * (1.0 - sig);
2954 }
2955 if scope.attn {
2956 add_vec_grad(
2957 block_grads.attn.w0.as_mut_slice(),
2958 scratch.grad_param.as_slice(),
2959 );
2960 add_outer_grad(
2961 block_grads.attn.w2.as_mut_slice(),
2962 c,
2963 d_w,
2964 scratch.grad_param.as_slice(),
2965 &tr.w_hidden.as_slice()[0..d_w],
2966 );
2967 }
2968 unsafe {
2969 kernel::gemv_t_avx(
2970 block.attn.w2.as_ptr(),
2971 scratch.grad_param.as_ptr(),
2972 scratch.grad_low_rank.as_mut_ptr(),
2973 c,
2974 d_w,
2975 );
2976 }
2977 for col in 0..d_w {
2978 let t = tr.w_hidden[col];
2979 scratch.grad_low_rank[col] *= 1.0 - t * t;
2980 }
2981 if scope.attn {
2982 add_outer_grad(
2983 block_grads.attn.w1.as_mut_slice(),
2984 d_w,
2985 c,
2986 &scratch.grad_low_rank.as_slice()[0..d_w],
2987 tr.xw.as_slice(),
2988 );
2989 }
2990 unsafe {
2991 kernel::gemv_t_avx(
2992 block.attn.w1.as_ptr(),
2993 scratch.grad_low_rank.as_ptr(),
2994 scratch.grad_x6.as_mut_ptr(),
2995 d_w,
2996 c,
2997 );
2998 }
2999
3000 for col in 0..c {
3001 let a = tr.a[col];
3002 scratch.grad_x5[col] *= a * (1.0 - a);
3003 }
3004 if scope.attn {
3005 add_vec_grad(
3006 block_grads.attn.a0.as_mut_slice(),
3007 scratch.grad_x5.as_slice(),
3008 );
3009 add_outer_grad(
3010 block_grads.attn.a2.as_mut_slice(),
3011 c,
3012 d_a,
3013 scratch.grad_x5.as_slice(),
3014 &tr.a_hidden.as_slice()[0..d_a],
3015 );
3016 }
3017 unsafe {
3018 kernel::gemv_t_avx(
3019 block.attn.a2.as_ptr(),
3020 scratch.grad_x5.as_ptr(),
3021 scratch.grad_low_rank.as_mut_ptr(),
3022 c,
3023 d_a,
3024 );
3025 }
3026 if scope.attn {
3027 add_outer_grad(
3028 block_grads.attn.a1.as_mut_slice(),
3029 d_a,
3030 c,
3031 &scratch.grad_low_rank.as_slice()[0..d_a],
3032 tr.xa.as_slice(),
3033 );
3034 }
3035 unsafe {
3036 kernel::gemv_t_avx(
3037 block.attn.a1.as_ptr(),
3038 scratch.grad_low_rank.as_ptr(),
3039 scratch.grad_x5.as_mut_ptr(),
3040 d_a,
3041 c,
3042 );
3043 }
3044
3045 if scope.attn {
3046 add_outer_grad(
3047 block_grads.attn.g2.as_mut_slice(),
3048 c,
3049 d_g,
3050 scratch.grad_saved.as_slice(),
3051 &tr.g_hidden.as_slice()[0..d_g],
3052 );
3053 }
3054 unsafe {
3055 kernel::gemv_t_avx(
3056 block.attn.g2.as_ptr(),
3057 scratch.grad_saved.as_ptr(),
3058 scratch.grad_low_rank.as_mut_ptr(),
3059 c,
3060 d_g,
3061 );
3062 }
3063 for col in 0..d_g {
3064 let sig = tr.g_hidden[col];
3065 scratch.grad_low_rank2[col] = scratch.grad_low_rank[col] * sig * (1.0 - sig);
3066 }
3067 if scope.attn {
3068 add_outer_grad(
3069 block_grads.attn.g1.as_mut_slice(),
3070 d_g,
3071 c,
3072 &scratch.grad_low_rank2.as_slice()[0..d_g],
3073 tr.xg.as_slice(),
3074 );
3075 }
3076 unsafe {
3077 kernel::gemv_t_avx(
3078 block.attn.g1.as_ptr(),
3079 scratch.grad_low_rank2.as_ptr(),
3080 scratch.grad_saved.as_mut_ptr(),
3081 d_g,
3082 c,
3083 );
3084 }
3085
3086 scratch
3087 .grad_x3
3088 .copy_from_slice(future_layer.att_x_prev.as_slice());
3089 future_layer.att_x_prev.zero();
3090
3091 for col in 0..c {
3092 let g = scratch.grad_param[col];
3093 let mix = block.attn.x_r[col];
3094 let base = tr.attn_norm[col];
3095 let prev = tr.att_x_prev_old[col];
3096 scratch.grad_x3[col] += g * (1.0 - mix);
3097 future_layer.att_x_prev[col] += g * mix;
3098 scratch.grad_x2[col] = g * (prev - base);
3099 }
3100 if scope.attn {
3101 add_vec_grad(
3102 block_grads.attn.x_r.as_mut_slice(),
3103 scratch.grad_x2.as_slice(),
3104 );
3105 }
3106
3107 for col in 0..c {
3108 let g = scratch.grad_x6[col];
3109 let mix = block.attn.x_w[col];
3110 let base = tr.attn_norm[col];
3111 let prev = tr.att_x_prev_old[col];
3112 scratch.grad_x3[col] += g * (1.0 - mix);
3113 future_layer.att_x_prev[col] += g * mix;
3114 scratch.grad_x2[col] = g * (prev - base);
3115 }
3116 if scope.attn {
3117 add_vec_grad(
3118 block_grads.attn.x_w.as_mut_slice(),
3119 scratch.grad_x2.as_slice(),
3120 );
3121 }
3122
3123 for col in 0..c {
3124 let g = scratch.grad_param2[col];
3125 let mix = block.attn.x_k[col];
3126 let base = tr.attn_norm[col];
3127 let prev = tr.att_x_prev_old[col];
3128 scratch.grad_x3[col] += g * (1.0 - mix);
3129 future_layer.att_x_prev[col] += g * mix;
3130 scratch.grad_x2[col] = g * (prev - base);
3131 }
3132 if scope.attn {
3133 add_vec_grad(
3134 block_grads.attn.x_k.as_mut_slice(),
3135 scratch.grad_x2.as_slice(),
3136 );
3137 }
3138
3139 for col in 0..c {
3140 let g = scratch.grad_x4[col];
3141 let mix = block.attn.x_v[col];
3142 let base = tr.attn_norm[col];
3143 let prev = tr.att_x_prev_old[col];
3144 scratch.grad_x3[col] += g * (1.0 - mix);
3145 future_layer.att_x_prev[col] += g * mix;
3146 scratch.grad_x2[col] = g * (prev - base);
3147 }
3148 if scope.attn {
3149 add_vec_grad(
3150 block_grads.attn.x_v.as_mut_slice(),
3151 scratch.grad_x2.as_slice(),
3152 );
3153 }
3154
3155 for col in 0..c {
3156 let g = scratch.grad_x5[col];
3157 let mix = block.attn.x_a[col];
3158 let base = tr.attn_norm[col];
3159 let prev = tr.att_x_prev_old[col];
3160 scratch.grad_x3[col] += g * (1.0 - mix);
3161 future_layer.att_x_prev[col] += g * mix;
3162 scratch.grad_x2[col] = g * (prev - base);
3163 }
3164 if scope.attn {
3165 add_vec_grad(
3166 block_grads.attn.x_a.as_mut_slice(),
3167 scratch.grad_x2.as_slice(),
3168 );
3169 }
3170
3171 for col in 0..c {
3172 let g = scratch.grad_saved[col];
3173 let mix = block.attn.x_g[col];
3174 let base = tr.attn_norm[col];
3175 let prev = tr.att_x_prev_old[col];
3176 scratch.grad_x3[col] += g * (1.0 - mix);
3177 future_layer.att_x_prev[col] += g * mix;
3178 scratch.grad_x2[col] = g * (prev - base);
3179 }
3180 if scope.attn {
3181 add_vec_grad(
3182 block_grads.attn.x_g.as_mut_slice(),
3183 scratch.grad_x2.as_slice(),
3184 );
3185 }
3186
3187 layer_norm_backward(
3188 tr.x_after_pre.as_slice(),
3189 block.attn_norm_w.as_slice(),
3190 scratch.grad_x3.as_slice(),
3191 self.cfg.layer_norm_eps,
3192 scratch.grad_x2.as_mut_slice(),
3193 scratch.grad_x4.as_mut_slice(),
3194 scratch.grad_x5.as_mut_slice(),
3195 );
3196 if scope.attn_norm {
3197 add_vec_grad(
3198 block_grads.attn_norm_w.as_mut_slice(),
3199 scratch.grad_x4.as_slice(),
3200 );
3201 add_vec_grad(
3202 block_grads.attn_norm_b.as_mut_slice(),
3203 scratch.grad_x5.as_slice(),
3204 );
3205 }
3206 for col in 0..c {
3207 scratch.grad_x[col] += scratch.grad_x2[col];
3208 }
3209
3210 if layer_idx == 0
3211 && let (Some(w), Some(_b)) = (&block.pre_norm_w, &block.pre_norm_b)
3212 {
3213 layer_norm_backward(
3214 tr.x_in.as_slice(),
3215 w.as_slice(),
3216 scratch.grad_x.as_slice(),
3217 self.cfg.layer_norm_eps,
3218 scratch.grad_x2.as_mut_slice(),
3219 scratch.grad_x3.as_mut_slice(),
3220 scratch.grad_x4.as_mut_slice(),
3221 );
3222 if scope.pre_norm {
3223 add_vec_grad(
3224 block_grads
3225 .pre_norm_w
3226 .as_mut()
3227 .expect("grad pre_norm_w")
3228 .as_mut_slice(),
3229 scratch.grad_x3.as_slice(),
3230 );
3231 add_vec_grad(
3232 block_grads
3233 .pre_norm_b
3234 .as_mut()
3235 .expect("grad pre_norm_b")
3236 .as_mut_slice(),
3237 scratch.grad_x4.as_slice(),
3238 );
3239 }
3240 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice());
3241 }
3242 }
3243
3244 if scope.embed {
3245 let token_idx = trace.token.min(self.cfg.vocab_size.saturating_sub(1));
3246 let off = token_idx * c;
3247 add_vec_grad(
3248 &mut grads.embeddings.as_mut_slice()[off..off + c],
3249 scratch.grad_x.as_slice(),
3250 );
3251 }
3252
3253 Ok(())
3254 }
3255
3256 #[allow(clippy::too_many_arguments)]
3257 pub(crate) fn online_train_segment_tbptt(
3259 &mut self,
3260 scratch: &mut ScratchBuffers,
3261 workspace: &mut TbpttReplayWorkspace,
3262 start_state: &State,
3263 steps: &[(u32, u8)],
3264 scope: TrainScopeMask,
3265 optimizer: OptimizerKind,
3266 lr: f32,
3267 clip: f32,
3268 replay_chunk: usize,
3269 adam_t: &mut usize,
3270 model_adam: Option<&mut FullAdamState>,
3271 out_bias: Option<&mut [f32]>,
3272 out_bias_adam_m: Option<&mut [f32]>,
3273 out_bias_adam_v: Option<&mut [f32]>,
3274 live_state_out: &mut State,
3275 ) -> Result<()> {
3276 if steps.is_empty() {
3277 live_state_out.copy_from(start_state);
3278 return Ok(());
3279 }
3280
3281 let grad_scale = 1.0f32 / (steps.len() as f32);
3282 let chunk = replay_chunk.max(1).min(steps.len().max(1));
3283 let TbpttReplayWorkspace {
3284 grads,
3285 recurrent,
3286 bias_grad: workspace_bias_grad,
3287 checkpoint_state,
3288 replay_state,
3289 checkpoints,
3290 step_states,
3291 step_traces,
3292 step_pdfs,
3293 } = workspace;
3294 grads.zero();
3295 recurrent.zero();
3296 let mut bias_grad = match out_bias.as_deref().map(<[f32]>::len) {
3297 Some(len) => {
3298 if workspace_bias_grad.len() != len {
3299 workspace_bias_grad.resize(len, 0.0);
3300 } else {
3301 workspace_bias_grad.fill(0.0);
3302 }
3303 Some(workspace_bias_grad.as_mut_slice())
3304 }
3305 None => None,
3306 };
3307
3308 {
3309 checkpoint_state.copy_from(start_state);
3310 let checkpoint_count = steps.len().div_ceil(chunk);
3311 ensure_cloned_len(checkpoints, checkpoint_count, start_state);
3312 checkpoints.truncate(checkpoint_count);
3313 scratch.set_capture_train_trace(false);
3314 for (checkpoint_idx, chunk_start) in (0..steps.len()).step_by(chunk).enumerate() {
3315 checkpoints[checkpoint_idx].copy_from(checkpoint_state);
3316 let chunk_end = (chunk_start + chunk).min(steps.len());
3317 for &(input_token, _) in &steps[chunk_start..chunk_end] {
3318 self.forward(scratch, input_token, checkpoint_state);
3319 }
3320 }
3321
3322 for chunk_idx in (0..checkpoint_count).rev() {
3323 let chunk_start = chunk_idx * chunk;
3324 let chunk_end = (chunk_start + chunk).min(steps.len());
3325 let chunk_steps = chunk_end - chunk_start;
3326 let checkpoint = &checkpoints[chunk_idx];
3327 replay_state.copy_from(checkpoint);
3328 let state_count = chunk_steps + 1;
3329 ensure_cloned_len(step_states, state_count, checkpoint);
3330 step_states.truncate(state_count);
3331 if step_traces.len() < chunk_steps {
3332 step_traces.resize_with(chunk_steps, || TokenTrainTrace::from_scratch(scratch));
3333 }
3334 step_traces.truncate(chunk_steps);
3335 let pdf_stride = self.cfg.vocab_size;
3336 step_pdfs.resize(chunk_steps.saturating_mul(pdf_stride), 0.0);
3337 step_states[0].copy_from(replay_state);
3338
3339 for (local_idx, &(input_token, _)) in
3340 steps[chunk_start..chunk_end].iter().enumerate()
3341 {
3342 scratch.set_capture_train_trace(true);
3343 let logits = self.forward(scratch, input_token, replay_state);
3344 let pdf_lo = local_idx * pdf_stride;
3345 let pdf_hi = pdf_lo + pdf_stride;
3346 super::super::softmax_pdf_floor_with_bias(
3347 logits,
3348 out_bias.as_deref(),
3349 &mut step_pdfs[pdf_lo..pdf_hi],
3350 );
3351 step_traces[local_idx].clone_from_scratch(scratch);
3352 step_states[local_idx + 1].copy_from(replay_state);
3353 }
3354
3355 for local_idx in (0..chunk_steps).rev() {
3356 let (_, target_symbol) = steps[chunk_start + local_idx];
3357 let pdf_lo = local_idx * pdf_stride;
3358 let pdf_hi = pdf_lo + pdf_stride;
3359 self.accumulate_token_step_gradients(
3360 scratch,
3361 &step_traces[local_idx],
3362 &step_states[local_idx + 1],
3363 target_symbol,
3364 &step_pdfs[pdf_lo..pdf_hi],
3365 grad_scale,
3366 scope,
3367 grads,
3368 bias_grad.as_deref_mut(),
3369 recurrent,
3370 )?;
3371 }
3372 }
3373 }
3374
3375 self.apply_full_gradients(
3376 grads,
3377 scope,
3378 optimizer,
3379 lr,
3380 clip,
3381 adam_t,
3382 model_adam,
3383 out_bias,
3384 bias_grad.as_deref(),
3385 out_bias_adam_m,
3386 out_bias_adam_v,
3387 )?;
3388
3389 scratch.set_capture_train_trace(false);
3390 live_state_out.copy_from(start_state);
3391 for &(input_token, _) in steps {
3392 self.forward(scratch, input_token, live_state_out);
3393 }
3394 Ok(())
3395 }
3396
3397 #[allow(clippy::too_many_arguments)]
3399 #[allow(clippy::needless_range_loop)]
3400 pub fn online_train_step_bptt1(
3401 &mut self,
3402 scratch: &mut ScratchBuffers,
3403 state: &State,
3404 symbol: u8,
3405 pdf: &[f64],
3406 scope: TrainScopeMask,
3407 optimizer: OptimizerKind,
3408 lr: f32,
3409 clip: f32,
3410 adam_t: &mut usize,
3411 model_adam: Option<&mut FullAdamState>,
3412 out_bias: Option<&mut [f32]>,
3413 out_bias_adam_m: Option<&mut [f32]>,
3414 out_bias_adam_v: Option<&mut [f32]>,
3415 ) -> Result<()> {
3416 if !scope.trains_any_params() {
3417 return Ok(());
3418 }
3419 if scope.trains_non_head_params() && !scratch.train_trace_valid {
3420 bail!("rwkv full training trace is missing; run one forward step first");
3421 }
3422 let c = self.cfg.hidden_size;
3423 let h = self.cfg.num_heads;
3424 let n = self.cfg.head_dim;
3425 let i = self.cfg.intermediate_size;
3426 let d_w = self.cfg.decay_low_rank;
3427 let d_a = self.cfg.a_low_rank;
3428 let d_v = self.cfg.v_low_rank;
3429 let d_g = self.cfg.g_low_rank;
3430 let vocab = self.cfg.vocab_size.min(pdf.len());
3431 if vocab == 0 {
3432 return Ok(());
3433 }
3434 let mut adam_step = None::<AdamStep>;
3435 let mut model_adam = model_adam;
3436 if matches!(optimizer, OptimizerKind::Adam) {
3437 *adam_t = adam_t.saturating_add(1);
3438 let t = (*adam_t).max(1) as i32;
3439 let b1 = 0.9f32;
3440 let b2 = 0.999f32;
3441 adam_step = Some(AdamStep {
3442 lr,
3443 clip: clip.max(0.0),
3444 b1,
3445 b2,
3446 eps: 1e-8,
3447 bias_corr1: 1.0 - b1.powi(t),
3448 bias_corr2: 1.0 - b2.powi(t),
3449 });
3450 if scope.trains_non_head_params() && model_adam.is_none() {
3451 bail!("rwkv Adam full-training state is missing");
3452 }
3453 }
3454
3455 scratch.grad_logits.zero();
3456 for idx in 0..vocab {
3457 let p = pdf[idx].clamp(1e-12, 1.0) as f32;
3458 let target = if idx == symbol as usize { 1.0 } else { 0.0 };
3459 let mut g = target - p;
3460 if clip > 0.0 {
3461 g = g.clamp(-clip, clip);
3462 }
3463 scratch.grad_logits[idx] = g;
3464 }
3465
3466 if scope.bias
3467 && let Some(bias) = out_bias
3468 {
3469 match optimizer {
3470 OptimizerKind::Sgd => {
3471 for idx in 0..bias.len().min(vocab) {
3472 bias[idx] += lr * scratch.grad_logits[idx];
3473 }
3474 }
3475 OptimizerKind::Adam => {
3476 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3477 let Some(m) = out_bias_adam_m else {
3478 bail!("rwkv Adam output-bias state is missing (m)");
3479 };
3480 let Some(vv) = out_bias_adam_v else {
3481 bail!("rwkv Adam output-bias state is missing (v)");
3482 };
3483 let n = bias.len().min(vocab);
3484 apply_adam_vec_update_raw(
3485 &mut bias[0..n],
3486 &scratch.grad_logits.as_slice()[0..n],
3487 &mut m[0..n],
3488 &mut vv[0..n],
3489 cfg,
3490 );
3491 }
3492 }
3493 }
3494
3495 scratch.grad_x.zero();
3496 if scope.head {
3497 match optimizer {
3498 OptimizerKind::Sgd => {
3499 fused_sgd_head_backward_update(
3500 self.lm_head.as_mut_slice(),
3501 vocab,
3502 c,
3503 &scratch.grad_logits.as_slice()[0..vocab],
3504 scratch.x_normed.as_slice(),
3505 scratch.grad_x.as_mut_slice(),
3506 lr,
3507 clip,
3508 );
3509 }
3510 OptimizerKind::Adam => {
3511 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3512 let adam = model_adam.as_mut().expect("adam state exists");
3513 fused_adam_head_backward_update(
3514 self.lm_head.as_mut_slice(),
3515 vocab,
3516 c,
3517 &scratch.grad_logits.as_slice()[0..vocab],
3518 scratch.x_normed.as_slice(),
3519 scratch.grad_x.as_mut_slice(),
3520 adam.lm_head.m.as_mut_slice(),
3521 adam.lm_head.v.as_mut_slice(),
3522 cfg,
3523 );
3524 }
3525 }
3526 } else {
3527 for row in 0..vocab {
3528 let g = scratch.grad_logits[row];
3529 if g == 0.0 {
3530 continue;
3531 }
3532 let row_off = row * c;
3533 for col in 0..c {
3534 scratch.grad_x[col] += self.lm_head[row_off + col] * g;
3535 }
3536 }
3537 }
3538
3539 let needs_backprop = scope.trains_non_head_params() || scope.head;
3540 if !needs_backprop {
3541 return Ok(());
3542 }
3543 layer_norm_backward(
3544 scratch.x.as_slice(),
3545 self.ln_out_w.as_slice(),
3546 scratch.grad_x.as_slice(),
3547 self.cfg.layer_norm_eps,
3548 scratch.grad_x2.as_mut_slice(),
3549 scratch.grad_x3.as_mut_slice(),
3550 scratch.grad_x4.as_mut_slice(),
3551 );
3552 if scope.head {
3553 match optimizer {
3554 OptimizerKind::Sgd => {
3555 sgd_vec_update(
3556 self.ln_out_w.as_mut_slice(),
3557 scratch.grad_x3.as_slice(),
3558 lr,
3559 clip,
3560 );
3561 sgd_vec_update(
3562 self.ln_out_b.as_mut_slice(),
3563 scratch.grad_x4.as_slice(),
3564 lr,
3565 clip,
3566 );
3567 }
3568 OptimizerKind::Adam => {
3569 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3570 let adam = model_adam.as_mut().expect("adam state exists");
3571 apply_adam_vec_update(
3572 self.ln_out_w.as_mut_slice(),
3573 scratch.grad_x3.as_slice(),
3574 &mut adam.ln_out_w,
3575 cfg,
3576 );
3577 apply_adam_vec_update(
3578 self.ln_out_b.as_mut_slice(),
3579 scratch.grad_x4.as_slice(),
3580 &mut adam.ln_out_b,
3581 cfg,
3582 );
3583 }
3584 }
3585 }
3586 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice());
3587 scratch.grad_v_first.zero();
3588
3589 for layer_idx in (0..self.cfg.num_layers).rev() {
3590 let tr = &scratch.train_trace_layers[layer_idx];
3591 let block = &mut self.blocks[layer_idx];
3592
3593 scratch.grad_x2.copy_from_slice(scratch.grad_x.as_slice()); scratch.grad_x3.copy_from_slice(scratch.grad_x.as_slice()); unsafe {
3599 kernel::gemv_t_avx(
3600 block.ffn.value_w.as_ptr(),
3601 scratch.grad_x3.as_ptr(),
3602 scratch.grad_ffn.as_mut_ptr(),
3603 c,
3604 i,
3605 );
3606 }
3607 if scope.ffn {
3608 match optimizer {
3609 OptimizerKind::Sgd => sgd_outer_update(
3610 block.ffn.value_w.as_mut_slice(),
3611 c,
3612 i,
3613 scratch.grad_x3.as_slice(),
3614 tr.ffn_k.as_slice(),
3615 lr,
3616 clip,
3617 ),
3618 OptimizerKind::Adam => {
3619 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3620 let adam =
3621 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3622 apply_adam_outer_update(
3623 block.ffn.value_w.as_mut_slice(),
3624 c,
3625 i,
3626 scratch.grad_x3.as_slice(),
3627 tr.ffn_k.as_slice(),
3628 &mut adam.ffn.value_w,
3629 cfg,
3630 );
3631 }
3632 }
3633 }
3634
3635 for col in 0..i {
3637 let pre = tr.ffn_pre[col];
3638 scratch.grad_ffn2[col] = if pre > 0.0 {
3639 scratch.grad_ffn[col] * (2.0 * pre)
3640 } else {
3641 0.0
3642 };
3643 }
3644
3645 unsafe {
3647 kernel::gemv_t_avx(
3648 block.ffn.key_w.as_ptr(),
3649 scratch.grad_ffn2.as_ptr(),
3650 scratch.grad_x4.as_mut_ptr(),
3651 i,
3652 c,
3653 );
3654 }
3655 if scope.ffn {
3656 match optimizer {
3657 OptimizerKind::Sgd => sgd_outer_update(
3658 block.ffn.key_w.as_mut_slice(),
3659 i,
3660 c,
3661 scratch.grad_ffn2.as_slice(),
3662 tr.ffn_xk.as_slice(),
3663 lr,
3664 clip,
3665 ),
3666 OptimizerKind::Adam => {
3667 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3668 let adam =
3669 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3670 apply_adam_outer_update(
3671 block.ffn.key_w.as_mut_slice(),
3672 i,
3673 c,
3674 scratch.grad_ffn2.as_slice(),
3675 tr.ffn_xk.as_slice(),
3676 &mut adam.ffn.key_w,
3677 cfg,
3678 );
3679 }
3680 }
3681 }
3682
3683 for col in 0..c {
3685 let g = scratch.grad_x4[col];
3686 let mix = block.ffn.x_k[col];
3687 let base = tr.ffn_norm[col];
3688 let prev = tr.ffn_x_prev_old[col];
3689 scratch.grad_x5[col] = g * (1.0 - mix); scratch.grad_param[col] = g * (prev - base); }
3692 if scope.ffn {
3693 match optimizer {
3694 OptimizerKind::Sgd => sgd_vec_update(
3695 block.ffn.x_k.as_mut_slice(),
3696 scratch.grad_param.as_slice(),
3697 lr,
3698 clip,
3699 ),
3700 OptimizerKind::Adam => {
3701 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3702 let adam =
3703 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3704 apply_adam_vec_update(
3705 block.ffn.x_k.as_mut_slice(),
3706 scratch.grad_param.as_slice(),
3707 &mut adam.ffn.x_k,
3708 cfg,
3709 );
3710 }
3711 }
3712 }
3713
3714 layer_norm_backward(
3716 tr.x_after_attn.as_slice(),
3717 block.ffn_norm_w.as_slice(),
3718 scratch.grad_x5.as_slice(),
3719 self.cfg.layer_norm_eps,
3720 scratch.grad_x4.as_mut_slice(),
3721 scratch.grad_x3.as_mut_slice(),
3722 scratch.grad_x6.as_mut_slice(),
3723 );
3724 if scope.ffn_norm {
3725 match optimizer {
3726 OptimizerKind::Sgd => {
3727 sgd_vec_update(
3728 block.ffn_norm_w.as_mut_slice(),
3729 scratch.grad_x3.as_slice(),
3730 lr,
3731 clip,
3732 );
3733 sgd_vec_update(
3734 block.ffn_norm_b.as_mut_slice(),
3735 scratch.grad_x6.as_slice(),
3736 lr,
3737 clip,
3738 );
3739 }
3740 OptimizerKind::Adam => {
3741 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3742 let adam =
3743 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3744 apply_adam_vec_update(
3745 block.ffn_norm_w.as_mut_slice(),
3746 scratch.grad_x3.as_slice(),
3747 &mut adam.ffn_norm_w,
3748 cfg,
3749 );
3750 apply_adam_vec_update(
3751 block.ffn_norm_b.as_mut_slice(),
3752 scratch.grad_x6.as_slice(),
3753 &mut adam.ffn_norm_b,
3754 cfg,
3755 );
3756 }
3757 }
3758 }
3759 for col in 0..c {
3760 scratch.grad_x2[col] += scratch.grad_x4[col];
3761 }
3762
3763 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice()); scratch.grad_x3.copy_from_slice(scratch.grad_x2.as_slice()); unsafe {
3769 kernel::gemv_t_avx(
3770 block.attn.o_proj.as_ptr(),
3771 scratch.grad_x3.as_ptr(),
3772 scratch.grad_x4.as_mut_ptr(),
3773 c,
3774 c,
3775 );
3776 }
3777 if scope.attn {
3778 match optimizer {
3779 OptimizerKind::Sgd => sgd_outer_update(
3780 block.attn.o_proj.as_mut_slice(),
3781 c,
3782 c,
3783 scratch.grad_x3.as_slice(),
3784 tr.y_gate.as_slice(),
3785 lr,
3786 clip,
3787 ),
3788 OptimizerKind::Adam => {
3789 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3790 let adam =
3791 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3792 apply_adam_outer_update(
3793 block.attn.o_proj.as_mut_slice(),
3794 c,
3795 c,
3796 scratch.grad_x3.as_slice(),
3797 tr.y_gate.as_slice(),
3798 &mut adam.attn.o_proj,
3799 cfg,
3800 );
3801 }
3802 }
3803 }
3804
3805 for col in 0..c {
3807 let gy = scratch.grad_x4[col];
3808 scratch.grad_saved[col] = gy * tr.y_head[col]; scratch.grad_x4[col] = gy * tr.g[col]; }
3811
3812 scratch.grad_x2.zero(); scratch.grad_x3.zero(); scratch.grad_x6.zero(); scratch.grad_param.zero(); for head_idx in 0..h {
3818 let off = head_idx * n;
3819 let mut g_alpha = 0.0f32;
3820 for j in 0..n {
3821 let g = scratch.grad_x4[off + j];
3822 g_alpha += g * tr.v[off + j];
3823 scratch.grad_x6[off + j] += g * tr.alpha[head_idx];
3824 }
3825 for j in 0..n {
3826 let idx = off + j;
3827 let rk = block.attn.r_k[idx];
3828 let rv = tr.r[idx];
3829 let kv = tr.k[idx];
3830 let g = g_alpha * rk;
3831 scratch.grad_x2[idx] += g * kv;
3832 scratch.grad_x3[idx] += g * rv;
3833 scratch.grad_param[idx] += g_alpha * rv * kv;
3834 }
3835 }
3836 if scope.attn {
3837 match optimizer {
3838 OptimizerKind::Sgd => sgd_vec_update(
3839 block.attn.r_k.as_mut_slice(),
3840 scratch.grad_param.as_slice(),
3841 lr,
3842 clip,
3843 ),
3844 OptimizerKind::Adam => {
3845 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3846 let adam =
3847 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3848 apply_adam_vec_update(
3849 block.attn.r_k.as_mut_slice(),
3850 scratch.grad_param.as_slice(),
3851 &mut adam.attn.r_k,
3852 cfg,
3853 );
3854 }
3855 }
3856 }
3857
3858 scratch.grad_x5.as_mut_slice()[0..c].copy_from_slice(&scratch.grad_x4.as_slice()[0..c]);
3860 group_norm_backward(
3861 tr.y_wkv.as_slice(),
3862 block.attn.g_norm_w.as_slice(),
3863 scratch.grad_x5.as_slice(),
3864 h,
3865 n,
3866 self.cfg.group_norm_eps,
3867 scratch.grad_x4.as_mut_slice(), scratch.grad_param.as_mut_slice(), scratch.grad_param2.as_mut_slice(), );
3871 if scope.attn {
3872 match optimizer {
3873 OptimizerKind::Sgd => {
3874 sgd_vec_update(
3875 block.attn.g_norm_w.as_mut_slice(),
3876 scratch.grad_param.as_slice(),
3877 lr,
3878 clip,
3879 );
3880 sgd_vec_update(
3881 block.attn.g_norm_b.as_mut_slice(),
3882 scratch.grad_param2.as_slice(),
3883 lr,
3884 clip,
3885 );
3886 }
3887 OptimizerKind::Adam => {
3888 let cfg = adam_step.as_ref().expect("adam cfg initialized");
3889 let adam =
3890 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
3891 apply_adam_vec_update(
3892 block.attn.g_norm_w.as_mut_slice(),
3893 scratch.grad_param.as_slice(),
3894 &mut adam.attn.g_norm_w,
3895 cfg,
3896 );
3897 apply_adam_vec_update(
3898 block.attn.g_norm_b.as_mut_slice(),
3899 scratch.grad_param2.as_slice(),
3900 &mut adam.attn.g_norm_b,
3901 cfg,
3902 );
3903 }
3904 }
3905 }
3906
3907 scratch.grad_param.zero(); scratch.grad_x5.zero(); scratch.grad_param2.zero(); let s_old = tr.att_state_old.as_slice();
3912 let s_new = state.layers[layer_idx].att_state.as_slice();
3913 for head_idx in 0..h {
3914 let off = head_idx * n;
3915 let s_head_old_off = head_idx * n * n;
3916 let s_head_new_off = head_idx * n * n;
3917 let grad_y = &scratch.grad_x4.as_slice()[off..off + n];
3918 let r_head = &tr.r.as_slice()[off..off + n];
3919 let k_head = &tr.k.as_slice()[off..off + n];
3920 let kk_head = &tr.kk.as_slice()[off..off + n];
3921 let a_head = &tr.a.as_slice()[off..off + n];
3922 let v_head = &tr.v.as_slice()[off..off + n];
3923
3924 unsafe {
3925 kernel::gemv_t_avx(
3926 s_new.as_ptr().add(s_head_new_off),
3927 grad_y.as_ptr(),
3928 scratch.grad_low_rank.as_mut_ptr(),
3929 n,
3930 n,
3931 );
3932 kernel::gemv_t_avx(
3933 s_old.as_ptr().add(s_head_old_off),
3934 grad_y.as_ptr(),
3935 scratch.grad_low_rank2.as_mut_ptr(),
3936 n,
3937 n,
3938 );
3939 }
3940
3941 for j in 0..n {
3942 let idx = off + j;
3943 scratch.grad_x2[idx] += scratch.grad_low_rank[j];
3944 scratch.grad_param[idx] += r_head[j] * scratch.grad_low_rank2[j];
3945 }
3946
3947 unsafe {
3948 kernel::gemv_avx(
3949 s_old.as_ptr().add(s_head_old_off),
3950 kk_head.as_ptr(),
3951 scratch.grad_low_rank.as_mut_ptr(),
3952 n,
3953 n,
3954 );
3955 }
3956
3957 let mut dot_gv = 0.0f32;
3958 let mut dot_rk = 0.0f32;
3959 let mut dot_r_kka = 0.0f32;
3960 let mut sum_gy_u = 0.0f32;
3961 for j in 0..n {
3962 dot_gv += grad_y[j] * v_head[j];
3963 dot_rk += r_head[j] * k_head[j];
3964 dot_r_kka += r_head[j] * kk_head[j] * a_head[j];
3965 sum_gy_u += grad_y[j] * scratch.grad_low_rank[j];
3966 }
3967
3968 for j in 0..n {
3969 let idx = off + j;
3970 scratch.grad_x3[idx] += r_head[j] * dot_gv;
3971 scratch.grad_x6[idx] += grad_y[j] * dot_rk;
3972 scratch.grad_x5[idx] -= sum_gy_u * r_head[j] * kk_head[j];
3973 scratch.grad_low_rank[j] = -grad_y[j] * dot_r_kka;
3974 }
3975
3976 unsafe {
3977 kernel::gemv_t_avx(
3978 s_old.as_ptr().add(s_head_old_off),
3979 scratch.grad_low_rank.as_ptr(),
3980 scratch.grad_low_rank2.as_mut_ptr(),
3981 n,
3982 n,
3983 );
3984 }
3985 for j in 0..n {
3986 let idx = off + j;
3987 scratch.grad_param2[idx] +=
3988 scratch.grad_low_rank2[j] - sum_gy_u * r_head[j] * a_head[j];
3989 }
3990 }
3991
3992 for col in 0..c {
3994 let gk = scratch.grad_x3[col];
3995 let scale = 1.0 + (tr.a[col] - 1.0) * block.attn.k_a[col];
3996 let d_scale = gk * tr.k_pre[col];
3997 scratch.grad_x3[col] = gk * scale; scratch.grad_x5[col] += d_scale * block.attn.k_a[col]; scratch.grad_param[col] = d_scale * (tr.a[col] - 1.0); }
4001 for head_idx in 0..h {
4002 let off = head_idx * n;
4003 l2_normalize_backward(
4004 &tr.kk_pre.as_slice()[off..off + n],
4005 &tr.kk.as_slice()[off..off + n],
4006 &scratch.grad_param2.as_slice()[off..off + n],
4007 1e-12,
4008 &mut scratch.grad_x4.as_mut_slice()[off..off + n],
4009 );
4010 }
4011 for col in 0..c {
4012 let g = scratch.grad_x4[col];
4013 scratch.grad_x3[col] += g * block.attn.k_k[col]; scratch.grad_param2[col] = g * tr.k_pre[col]; }
4016 if scope.attn {
4017 match optimizer {
4018 OptimizerKind::Sgd => {
4019 sgd_vec_update(
4020 block.attn.k_a.as_mut_slice(),
4021 scratch.grad_param.as_slice(),
4022 lr,
4023 clip,
4024 );
4025 sgd_vec_update(
4026 block.attn.k_k.as_mut_slice(),
4027 scratch.grad_param2.as_slice(),
4028 lr,
4029 clip,
4030 );
4031 }
4032 OptimizerKind::Adam => {
4033 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4034 let adam =
4035 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4036 apply_adam_vec_update(
4037 block.attn.k_a.as_mut_slice(),
4038 scratch.grad_param.as_slice(),
4039 &mut adam.attn.k_a,
4040 cfg,
4041 );
4042 apply_adam_vec_update(
4043 block.attn.k_k.as_mut_slice(),
4044 scratch.grad_param2.as_slice(),
4045 &mut adam.attn.k_k,
4046 cfg,
4047 );
4048 }
4049 }
4050 }
4051
4052 scratch
4054 .grad_param2
4055 .copy_from_slice(scratch.grad_x6.as_slice()); if layer_idx == 0 {
4057 for col in 0..c {
4058 scratch.grad_x6[col] += scratch.grad_v_first[col];
4059 }
4060 } else if tr.uses_v_residual
4061 && let (Some(v1), Some(v2), Some(v0)) =
4062 (&mut block.attn.v1, &mut block.attn.v2, &mut block.attn.v0)
4063 {
4064 for col in 0..c {
4065 let gv = scratch.grad_param2[col];
4066 let nu = tr.nu[col];
4067 scratch.grad_x6[col] = gv * (1.0 - nu); scratch.grad_x3[col] = gv * (scratch.train_v_first[col] - tr.v_pre[col]); scratch.grad_v_first[col] += gv * nu; }
4071 for col in 0..c {
4072 let nu = tr.nu[col];
4073 scratch.grad_x3[col] *= nu * (1.0 - nu); }
4075 if scope.attn {
4076 match optimizer {
4077 OptimizerKind::Sgd => {
4078 sgd_vec_update(v0.as_mut_slice(), scratch.grad_x3.as_slice(), lr, clip)
4079 }
4080 OptimizerKind::Adam => {
4081 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4082 let adam = &mut model_adam.as_mut().expect("adam state exists").blocks
4083 [layer_idx];
4084 apply_adam_vec_update(
4085 v0.as_mut_slice(),
4086 scratch.grad_x3.as_slice(),
4087 adam.attn.v0.as_mut().expect("adam v0 state"),
4088 cfg,
4089 );
4090 }
4091 }
4092 }
4093 if scope.attn {
4094 match optimizer {
4095 OptimizerKind::Sgd => sgd_outer_update(
4096 v2.as_mut_slice(),
4097 c,
4098 d_v,
4099 scratch.grad_x3.as_slice(),
4100 &tr.v_hidden.as_slice()[0..d_v],
4101 lr,
4102 clip,
4103 ),
4104 OptimizerKind::Adam => {
4105 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4106 let adam = &mut model_adam.as_mut().expect("adam state exists").blocks
4107 [layer_idx];
4108 apply_adam_outer_update(
4109 v2.as_mut_slice(),
4110 c,
4111 d_v,
4112 scratch.grad_x3.as_slice(),
4113 &tr.v_hidden.as_slice()[0..d_v],
4114 adam.attn.v2.as_mut().expect("adam v2 state"),
4115 cfg,
4116 );
4117 }
4118 }
4119 }
4120 unsafe {
4121 kernel::gemv_t_avx(
4122 v2.as_ptr(),
4123 scratch.grad_x3.as_ptr(),
4124 scratch.grad_low_rank.as_mut_ptr(),
4125 c,
4126 d_v,
4127 );
4128 }
4129 if scope.attn {
4130 match optimizer {
4131 OptimizerKind::Sgd => sgd_outer_update(
4132 v1.as_mut_slice(),
4133 d_v,
4134 c,
4135 &scratch.grad_low_rank.as_slice()[0..d_v],
4136 tr.xv.as_slice(),
4137 lr,
4138 clip,
4139 ),
4140 OptimizerKind::Adam => {
4141 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4142 let adam = &mut model_adam.as_mut().expect("adam state exists").blocks
4143 [layer_idx];
4144 apply_adam_outer_update(
4145 v1.as_mut_slice(),
4146 d_v,
4147 c,
4148 &scratch.grad_low_rank.as_slice()[0..d_v],
4149 tr.xv.as_slice(),
4150 adam.attn.v1.as_mut().expect("adam v1 state"),
4151 cfg,
4152 );
4153 }
4154 }
4155 }
4156 for col in 0..c {
4157 let mut acc = 0.0f32;
4158 for row in 0..d_v {
4159 acc += v1[row * c + col] * scratch.grad_low_rank[row];
4160 }
4161 scratch.grad_x4[col] += acc; }
4163 }
4164
4165 let proj_size = c * c;
4167 if scope.attn {
4168 match optimizer {
4169 OptimizerKind::Sgd => {
4170 sgd_outer_update(
4171 &mut block.attn.rkv_proj.as_mut_slice()[0..proj_size],
4172 c,
4173 c,
4174 scratch.grad_x2.as_slice(),
4175 tr.xr.as_slice(),
4176 lr,
4177 clip,
4178 );
4179 sgd_outer_update(
4180 &mut block.attn.rkv_proj.as_mut_slice()[proj_size..2 * proj_size],
4181 c,
4182 c,
4183 scratch.grad_x3.as_slice(),
4184 tr.xk.as_slice(),
4185 lr,
4186 clip,
4187 );
4188 sgd_outer_update(
4189 &mut block.attn.rkv_proj.as_mut_slice()[2 * proj_size..3 * proj_size],
4190 c,
4191 c,
4192 scratch.grad_x6.as_slice(),
4193 tr.xv.as_slice(),
4194 lr,
4195 clip,
4196 );
4197 }
4198 OptimizerKind::Adam => {
4199 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4200 let adam =
4201 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4202 apply_adam_outer_update_raw(
4203 &mut block.attn.rkv_proj.as_mut_slice()[0..proj_size],
4204 c,
4205 c,
4206 scratch.grad_x2.as_slice(),
4207 tr.xr.as_slice(),
4208 &mut adam.attn.rkv_proj.m.as_mut_slice()[0..proj_size],
4209 &mut adam.attn.rkv_proj.v.as_mut_slice()[0..proj_size],
4210 cfg,
4211 );
4212 apply_adam_outer_update_raw(
4213 &mut block.attn.rkv_proj.as_mut_slice()[proj_size..2 * proj_size],
4214 c,
4215 c,
4216 scratch.grad_x3.as_slice(),
4217 tr.xk.as_slice(),
4218 &mut adam.attn.rkv_proj.m.as_mut_slice()[proj_size..2 * proj_size],
4219 &mut adam.attn.rkv_proj.v.as_mut_slice()[proj_size..2 * proj_size],
4220 cfg,
4221 );
4222 apply_adam_outer_update_raw(
4223 &mut block.attn.rkv_proj.as_mut_slice()[2 * proj_size..3 * proj_size],
4224 c,
4225 c,
4226 scratch.grad_x6.as_slice(),
4227 tr.xv.as_slice(),
4228 &mut adam.attn.rkv_proj.m.as_mut_slice()[2 * proj_size..3 * proj_size],
4229 &mut adam.attn.rkv_proj.v.as_mut_slice()[2 * proj_size..3 * proj_size],
4230 cfg,
4231 );
4232 }
4233 }
4234 }
4235 let proj = block.attn.rkv_proj.as_slice();
4236 unsafe {
4237 kernel::gemv_t_avx(
4238 proj.as_ptr(),
4239 scratch.grad_x2.as_ptr(),
4240 scratch.grad_param.as_mut_ptr(),
4241 c,
4242 c,
4243 );
4244 kernel::gemv_t_avx(
4245 proj.as_ptr().add(proj_size),
4246 scratch.grad_x3.as_ptr(),
4247 scratch.grad_param2.as_mut_ptr(),
4248 c,
4249 c,
4250 );
4251 kernel::gemv_t_avx(
4252 proj.as_ptr().add(2 * proj_size),
4253 scratch.grad_x6.as_ptr(),
4254 scratch.grad_x4.as_mut_ptr(),
4255 c,
4256 c,
4257 );
4258 }
4259
4260 let inv_sqrt_e = 1.0 / std::f32::consts::E.sqrt();
4262 for col in 0..c {
4263 let sig = tr.w_sigmoid[col];
4264 let d_sig = scratch.grad_param[col] * (-inv_sqrt_e) * tr.w_decay[col];
4265 scratch.grad_param[col] = d_sig * sig * (1.0 - sig); }
4267 if scope.attn {
4268 match optimizer {
4269 OptimizerKind::Sgd => sgd_vec_update(
4270 block.attn.w0.as_mut_slice(),
4271 scratch.grad_param.as_slice(),
4272 lr,
4273 clip,
4274 ),
4275 OptimizerKind::Adam => {
4276 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4277 let adam =
4278 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4279 apply_adam_vec_update(
4280 block.attn.w0.as_mut_slice(),
4281 scratch.grad_param.as_slice(),
4282 &mut adam.attn.w0,
4283 cfg,
4284 );
4285 }
4286 }
4287 match optimizer {
4288 OptimizerKind::Sgd => sgd_outer_update(
4289 block.attn.w2.as_mut_slice(),
4290 c,
4291 d_w,
4292 scratch.grad_param.as_slice(),
4293 &tr.w_hidden.as_slice()[0..d_w],
4294 lr,
4295 clip,
4296 ),
4297 OptimizerKind::Adam => {
4298 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4299 let adam =
4300 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4301 apply_adam_outer_update(
4302 block.attn.w2.as_mut_slice(),
4303 c,
4304 d_w,
4305 scratch.grad_param.as_slice(),
4306 &tr.w_hidden.as_slice()[0..d_w],
4307 &mut adam.attn.w2,
4308 cfg,
4309 );
4310 }
4311 }
4312 }
4313 unsafe {
4314 kernel::gemv_t_avx(
4315 block.attn.w2.as_ptr(),
4316 scratch.grad_param.as_ptr(),
4317 scratch.grad_low_rank.as_mut_ptr(),
4318 c,
4319 d_w,
4320 );
4321 }
4322 for col in 0..d_w {
4323 let t = tr.w_hidden[col];
4324 scratch.grad_low_rank[col] *= 1.0 - t * t;
4325 }
4326 if scope.attn {
4327 match optimizer {
4328 OptimizerKind::Sgd => sgd_outer_update(
4329 block.attn.w1.as_mut_slice(),
4330 d_w,
4331 c,
4332 &scratch.grad_low_rank.as_slice()[0..d_w],
4333 tr.xw.as_slice(),
4334 lr,
4335 clip,
4336 ),
4337 OptimizerKind::Adam => {
4338 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4339 let adam =
4340 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4341 apply_adam_outer_update(
4342 block.attn.w1.as_mut_slice(),
4343 d_w,
4344 c,
4345 &scratch.grad_low_rank.as_slice()[0..d_w],
4346 tr.xw.as_slice(),
4347 &mut adam.attn.w1,
4348 cfg,
4349 );
4350 }
4351 }
4352 }
4353 unsafe {
4354 kernel::gemv_t_avx(
4355 block.attn.w1.as_ptr(),
4356 scratch.grad_low_rank.as_ptr(),
4357 scratch.grad_x6.as_mut_ptr(),
4358 d_w,
4359 c,
4360 );
4361 }
4362
4363 for col in 0..c {
4365 let a = tr.a[col];
4366 scratch.grad_x5[col] *= a * (1.0 - a); }
4368 if scope.attn {
4369 match optimizer {
4370 OptimizerKind::Sgd => sgd_vec_update(
4371 block.attn.a0.as_mut_slice(),
4372 scratch.grad_x5.as_slice(),
4373 lr,
4374 clip,
4375 ),
4376 OptimizerKind::Adam => {
4377 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4378 let adam =
4379 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4380 apply_adam_vec_update(
4381 block.attn.a0.as_mut_slice(),
4382 scratch.grad_x5.as_slice(),
4383 &mut adam.attn.a0,
4384 cfg,
4385 );
4386 }
4387 }
4388 match optimizer {
4389 OptimizerKind::Sgd => sgd_outer_update(
4390 block.attn.a2.as_mut_slice(),
4391 c,
4392 d_a,
4393 scratch.grad_x5.as_slice(),
4394 &tr.a_hidden.as_slice()[0..d_a],
4395 lr,
4396 clip,
4397 ),
4398 OptimizerKind::Adam => {
4399 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4400 let adam =
4401 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4402 apply_adam_outer_update(
4403 block.attn.a2.as_mut_slice(),
4404 c,
4405 d_a,
4406 scratch.grad_x5.as_slice(),
4407 &tr.a_hidden.as_slice()[0..d_a],
4408 &mut adam.attn.a2,
4409 cfg,
4410 );
4411 }
4412 }
4413 }
4414 unsafe {
4415 kernel::gemv_t_avx(
4416 block.attn.a2.as_ptr(),
4417 scratch.grad_x5.as_ptr(),
4418 scratch.grad_low_rank.as_mut_ptr(),
4419 c,
4420 d_a,
4421 );
4422 }
4423 if scope.attn {
4424 match optimizer {
4425 OptimizerKind::Sgd => sgd_outer_update(
4426 block.attn.a1.as_mut_slice(),
4427 d_a,
4428 c,
4429 &scratch.grad_low_rank.as_slice()[0..d_a],
4430 tr.xa.as_slice(),
4431 lr,
4432 clip,
4433 ),
4434 OptimizerKind::Adam => {
4435 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4436 let adam =
4437 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4438 apply_adam_outer_update(
4439 block.attn.a1.as_mut_slice(),
4440 d_a,
4441 c,
4442 &scratch.grad_low_rank.as_slice()[0..d_a],
4443 tr.xa.as_slice(),
4444 &mut adam.attn.a1,
4445 cfg,
4446 );
4447 }
4448 }
4449 }
4450 unsafe {
4451 kernel::gemv_t_avx(
4452 block.attn.a1.as_ptr(),
4453 scratch.grad_low_rank.as_ptr(),
4454 scratch.grad_x5.as_mut_ptr(),
4455 d_a,
4456 c,
4457 );
4458 }
4459
4460 if scope.attn {
4462 match optimizer {
4463 OptimizerKind::Sgd => sgd_outer_update(
4464 block.attn.g2.as_mut_slice(),
4465 c,
4466 d_g,
4467 scratch.grad_saved.as_slice(),
4468 &tr.g_hidden.as_slice()[0..d_g],
4469 lr,
4470 clip,
4471 ),
4472 OptimizerKind::Adam => {
4473 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4474 let adam =
4475 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4476 apply_adam_outer_update(
4477 block.attn.g2.as_mut_slice(),
4478 c,
4479 d_g,
4480 scratch.grad_saved.as_slice(),
4481 &tr.g_hidden.as_slice()[0..d_g],
4482 &mut adam.attn.g2,
4483 cfg,
4484 );
4485 }
4486 }
4487 }
4488 unsafe {
4489 kernel::gemv_t_avx(
4490 block.attn.g2.as_ptr(),
4491 scratch.grad_saved.as_ptr(),
4492 scratch.grad_low_rank.as_mut_ptr(),
4493 c,
4494 d_g,
4495 );
4496 }
4497 for col in 0..d_g {
4498 let sig = tr.g_hidden[col];
4499 scratch.grad_low_rank2[col] = scratch.grad_low_rank[col] * sig * (1.0 - sig);
4500 }
4501 if scope.attn {
4502 match optimizer {
4503 OptimizerKind::Sgd => sgd_outer_update(
4504 block.attn.g1.as_mut_slice(),
4505 d_g,
4506 c,
4507 &scratch.grad_low_rank2.as_slice()[0..d_g],
4508 tr.xg.as_slice(),
4509 lr,
4510 clip,
4511 ),
4512 OptimizerKind::Adam => {
4513 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4514 let adam =
4515 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4516 apply_adam_outer_update(
4517 block.attn.g1.as_mut_slice(),
4518 d_g,
4519 c,
4520 &scratch.grad_low_rank2.as_slice()[0..d_g],
4521 tr.xg.as_slice(),
4522 &mut adam.attn.g1,
4523 cfg,
4524 );
4525 }
4526 }
4527 }
4528 unsafe {
4529 kernel::gemv_t_avx(
4530 block.attn.g1.as_ptr(),
4531 scratch.grad_low_rank2.as_ptr(),
4532 scratch.grad_saved.as_mut_ptr(),
4533 d_g,
4534 c,
4535 );
4536 }
4537
4538 scratch.grad_x3.zero(); for col in 0..c {
4543 let g = scratch.grad_param[col];
4544 let mix = block.attn.x_r[col];
4545 let base = tr.attn_norm[col];
4546 let prev = tr.att_x_prev_old[col];
4547 scratch.grad_x3[col] += g * (1.0 - mix);
4548 scratch.grad_x2[col] = g * (prev - base);
4549 }
4550 if scope.attn {
4551 match optimizer {
4552 OptimizerKind::Sgd => sgd_vec_update(
4553 block.attn.x_r.as_mut_slice(),
4554 scratch.grad_x2.as_slice(),
4555 lr,
4556 clip,
4557 ),
4558 OptimizerKind::Adam => {
4559 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4560 let adam =
4561 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4562 apply_adam_vec_update(
4563 block.attn.x_r.as_mut_slice(),
4564 scratch.grad_x2.as_slice(),
4565 &mut adam.attn.x_r,
4566 cfg,
4567 );
4568 }
4569 }
4570 }
4571
4572 for col in 0..c {
4574 let g = scratch.grad_x6[col];
4575 let mix = block.attn.x_w[col];
4576 let base = tr.attn_norm[col];
4577 let prev = tr.att_x_prev_old[col];
4578 scratch.grad_x3[col] += g * (1.0 - mix);
4579 scratch.grad_x2[col] = g * (prev - base);
4580 }
4581 if scope.attn {
4582 match optimizer {
4583 OptimizerKind::Sgd => sgd_vec_update(
4584 block.attn.x_w.as_mut_slice(),
4585 scratch.grad_x2.as_slice(),
4586 lr,
4587 clip,
4588 ),
4589 OptimizerKind::Adam => {
4590 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4591 let adam =
4592 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4593 apply_adam_vec_update(
4594 block.attn.x_w.as_mut_slice(),
4595 scratch.grad_x2.as_slice(),
4596 &mut adam.attn.x_w,
4597 cfg,
4598 );
4599 }
4600 }
4601 }
4602
4603 for col in 0..c {
4605 let g = scratch.grad_param2[col];
4606 let mix = block.attn.x_k[col];
4607 let base = tr.attn_norm[col];
4608 let prev = tr.att_x_prev_old[col];
4609 scratch.grad_x3[col] += g * (1.0 - mix);
4610 scratch.grad_x2[col] = g * (prev - base);
4611 }
4612 if scope.attn {
4613 match optimizer {
4614 OptimizerKind::Sgd => sgd_vec_update(
4615 block.attn.x_k.as_mut_slice(),
4616 scratch.grad_x2.as_slice(),
4617 lr,
4618 clip,
4619 ),
4620 OptimizerKind::Adam => {
4621 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4622 let adam =
4623 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4624 apply_adam_vec_update(
4625 block.attn.x_k.as_mut_slice(),
4626 scratch.grad_x2.as_slice(),
4627 &mut adam.attn.x_k,
4628 cfg,
4629 );
4630 }
4631 }
4632 }
4633
4634 for col in 0..c {
4636 let g = scratch.grad_x4[col];
4637 let mix = block.attn.x_v[col];
4638 let base = tr.attn_norm[col];
4639 let prev = tr.att_x_prev_old[col];
4640 scratch.grad_x3[col] += g * (1.0 - mix);
4641 scratch.grad_x2[col] = g * (prev - base);
4642 }
4643 if scope.attn {
4644 match optimizer {
4645 OptimizerKind::Sgd => sgd_vec_update(
4646 block.attn.x_v.as_mut_slice(),
4647 scratch.grad_x2.as_slice(),
4648 lr,
4649 clip,
4650 ),
4651 OptimizerKind::Adam => {
4652 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4653 let adam =
4654 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4655 apply_adam_vec_update(
4656 block.attn.x_v.as_mut_slice(),
4657 scratch.grad_x2.as_slice(),
4658 &mut adam.attn.x_v,
4659 cfg,
4660 );
4661 }
4662 }
4663 }
4664
4665 for col in 0..c {
4667 let g = scratch.grad_x5[col];
4668 let mix = block.attn.x_a[col];
4669 let base = tr.attn_norm[col];
4670 let prev = tr.att_x_prev_old[col];
4671 scratch.grad_x3[col] += g * (1.0 - mix);
4672 scratch.grad_x2[col] = g * (prev - base);
4673 }
4674 if scope.attn {
4675 match optimizer {
4676 OptimizerKind::Sgd => sgd_vec_update(
4677 block.attn.x_a.as_mut_slice(),
4678 scratch.grad_x2.as_slice(),
4679 lr,
4680 clip,
4681 ),
4682 OptimizerKind::Adam => {
4683 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4684 let adam =
4685 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4686 apply_adam_vec_update(
4687 block.attn.x_a.as_mut_slice(),
4688 scratch.grad_x2.as_slice(),
4689 &mut adam.attn.x_a,
4690 cfg,
4691 );
4692 }
4693 }
4694 }
4695
4696 for col in 0..c {
4698 let g = scratch.grad_saved[col];
4699 let mix = block.attn.x_g[col];
4700 let base = tr.attn_norm[col];
4701 let prev = tr.att_x_prev_old[col];
4702 scratch.grad_x3[col] += g * (1.0 - mix);
4703 scratch.grad_x2[col] = g * (prev - base);
4704 }
4705 if scope.attn {
4706 match optimizer {
4707 OptimizerKind::Sgd => sgd_vec_update(
4708 block.attn.x_g.as_mut_slice(),
4709 scratch.grad_x2.as_slice(),
4710 lr,
4711 clip,
4712 ),
4713 OptimizerKind::Adam => {
4714 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4715 let adam =
4716 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4717 apply_adam_vec_update(
4718 block.attn.x_g.as_mut_slice(),
4719 scratch.grad_x2.as_slice(),
4720 &mut adam.attn.x_g,
4721 cfg,
4722 );
4723 }
4724 }
4725 }
4726
4727 layer_norm_backward(
4729 tr.x_after_pre.as_slice(),
4730 block.attn_norm_w.as_slice(),
4731 scratch.grad_x3.as_slice(),
4732 self.cfg.layer_norm_eps,
4733 scratch.grad_x2.as_mut_slice(),
4734 scratch.grad_x4.as_mut_slice(),
4735 scratch.grad_x5.as_mut_slice(),
4736 );
4737 if scope.attn_norm {
4738 match optimizer {
4739 OptimizerKind::Sgd => {
4740 sgd_vec_update(
4741 block.attn_norm_w.as_mut_slice(),
4742 scratch.grad_x4.as_slice(),
4743 lr,
4744 clip,
4745 );
4746 sgd_vec_update(
4747 block.attn_norm_b.as_mut_slice(),
4748 scratch.grad_x5.as_slice(),
4749 lr,
4750 clip,
4751 );
4752 }
4753 OptimizerKind::Adam => {
4754 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4755 let adam =
4756 &mut model_adam.as_mut().expect("adam state exists").blocks[layer_idx];
4757 apply_adam_vec_update(
4758 block.attn_norm_w.as_mut_slice(),
4759 scratch.grad_x4.as_slice(),
4760 &mut adam.attn_norm_w,
4761 cfg,
4762 );
4763 apply_adam_vec_update(
4764 block.attn_norm_b.as_mut_slice(),
4765 scratch.grad_x5.as_slice(),
4766 &mut adam.attn_norm_b,
4767 cfg,
4768 );
4769 }
4770 }
4771 }
4772 for col in 0..c {
4773 scratch.grad_x[col] += scratch.grad_x2[col];
4774 }
4775
4776 if layer_idx == 0
4778 && let (Some(w), Some(b)) = (&mut block.pre_norm_w, &mut block.pre_norm_b)
4779 {
4780 layer_norm_backward(
4781 tr.x_in.as_slice(),
4782 w.as_slice(),
4783 scratch.grad_x.as_slice(),
4784 self.cfg.layer_norm_eps,
4785 scratch.grad_x2.as_mut_slice(),
4786 scratch.grad_x3.as_mut_slice(),
4787 scratch.grad_x4.as_mut_slice(),
4788 );
4789 if scope.pre_norm {
4790 match optimizer {
4791 OptimizerKind::Sgd => {
4792 sgd_vec_update(w.as_mut_slice(), scratch.grad_x3.as_slice(), lr, clip);
4793 sgd_vec_update(b.as_mut_slice(), scratch.grad_x4.as_slice(), lr, clip);
4794 }
4795 OptimizerKind::Adam => {
4796 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4797 let adam = &mut model_adam.as_mut().expect("adam state exists").blocks
4798 [layer_idx];
4799 apply_adam_vec_update(
4800 w.as_mut_slice(),
4801 scratch.grad_x3.as_slice(),
4802 adam.pre_norm_w.as_mut().expect("adam pre_norm_w"),
4803 cfg,
4804 );
4805 apply_adam_vec_update(
4806 b.as_mut_slice(),
4807 scratch.grad_x4.as_slice(),
4808 adam.pre_norm_b.as_mut().expect("adam pre_norm_b"),
4809 cfg,
4810 );
4811 }
4812 }
4813 }
4814 scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice());
4815 }
4816 }
4817
4818 if scope.embed {
4819 let token_idx = scratch
4820 .train_token
4821 .min(self.cfg.vocab_size.saturating_sub(1));
4822 let off = token_idx * c;
4823 let row = &mut self.embeddings.as_mut_slice()[off..off + c];
4824 match optimizer {
4825 OptimizerKind::Sgd => {
4826 sgd_vec_update(row, scratch.grad_x.as_slice(), lr, clip);
4827 }
4828 OptimizerKind::Adam => {
4829 let cfg = adam_step.as_ref().expect("adam cfg initialized");
4830 let adam = model_adam.as_mut().expect("adam state exists");
4831 let m = &mut adam.embeddings.m.as_mut_slice()[off..off + c];
4832 let v = &mut adam.embeddings.v.as_mut_slice()[off..off + c];
4833 apply_adam_vec_update_raw(row, scratch.grad_x.as_slice(), m, v, cfg);
4834 }
4835 }
4836 }
4837 Ok(())
4838 }
4839
4840 #[inline(never)]
4843 pub fn forward<'a>(
4844 &'a self,
4845 scratch: &'a mut ScratchBuffers,
4846 token: u32,
4847 state: &mut State,
4848 ) -> &'a [f32] {
4849 let mut sink = NullProfiler;
4850 self.forward_with_sink(scratch, token, state, &mut sink)
4851 }
4852
4853 #[inline(never)]
4855 pub fn forward_with_profiler<'a, S: ProfilerSink>(
4856 &'a self,
4857 scratch: &'a mut ScratchBuffers,
4858 token: u32,
4859 state: &mut State,
4860 profiler: &mut S,
4861 ) -> &'a [f32] {
4862 self.forward_with_sink(scratch, token, state, profiler)
4863 }
4864
4865 #[inline(never)]
4866 fn forward_with_sink<'a, S: ProfilerSink>(
4867 &'a self,
4868 scratch: &'a mut ScratchBuffers,
4869 token: u32,
4870 state: &mut State,
4871 profiler: &mut S,
4872 ) -> &'a [f32] {
4873 if scratch.capture_train_trace {
4874 self.forward_with_sink_impl::<true, S>(scratch, token, state, profiler)
4875 } else {
4876 self.forward_with_sink_impl::<false, S>(scratch, token, state, profiler)
4877 }
4878 }
4879
4880 fn forward_with_sink_impl<'a, const CAPTURE: bool, S: ProfilerSink>(
4881 &'a self,
4882 scratch: &'a mut ScratchBuffers,
4883 token: u32,
4884 state: &mut State,
4885 profiler: &mut S,
4886 ) -> &'a [f32] {
4887 let c = self.cfg.hidden_size;
4888 let _h = self.cfg.num_heads;
4889 let _n = self.cfg.head_dim;
4890 let num_layers = self.cfg.num_layers;
4891 let token_idx = (token as usize).min(self.cfg.vocab_size.saturating_sub(1));
4892
4893 let emb_offset = token_idx * c;
4895 let emb_slice = &self.embeddings.as_slice()[emb_offset..emb_offset + c];
4896 scratch.x.as_mut_slice().copy_from_slice(emb_slice);
4897 if CAPTURE {
4898 scratch.train_token = token_idx;
4899 scratch.train_trace_valid = true;
4900 } else {
4901 scratch.train_trace_valid = false;
4902 }
4903
4904 profiler.begin_token();
4905
4906 unsafe {
4907 for layer_idx in 0..num_layers {
4909 if CAPTURE {
4910 scratch.train_trace_layers[layer_idx]
4911 .x_in
4912 .copy_from(&scratch.x);
4913 }
4914 if let (Some(w), Some(b)) = (
4916 &self.blocks[layer_idx].pre_norm_w,
4917 &self.blocks[layer_idx].pre_norm_b,
4918 ) {
4919 kernel::layer_norm_avx(
4920 scratch.x.as_ptr(),
4921 w.as_ptr(),
4922 b.as_ptr(),
4923 scratch.x.as_mut_ptr(),
4924 c,
4925 self.cfg.layer_norm_eps,
4926 );
4927 }
4928 if CAPTURE {
4929 scratch.train_trace_layers[layer_idx]
4930 .x_after_pre
4931 .copy_from(&scratch.x);
4932 }
4933
4934 kernel::layer_norm_avx(
4936 scratch.x.as_ptr(),
4937 self.blocks[layer_idx].attn_norm_w.as_ptr(),
4938 self.blocks[layer_idx].attn_norm_b.as_ptr(),
4939 scratch.x_normed.as_mut_ptr(),
4940 c,
4941 self.cfg.layer_norm_eps,
4942 );
4943 if CAPTURE {
4944 scratch.train_trace_layers[layer_idx]
4945 .attn_norm
4946 .copy_from(&scratch.x_normed);
4947 }
4948
4949 let trace_ptr = if CAPTURE {
4950 &mut scratch.train_trace_layers[layer_idx] as *mut LayerTrainTrace
4951 } else {
4952 std::ptr::null_mut()
4953 };
4954 if S::ENABLED {
4955 let attn_start = Instant::now();
4956 self.attention_forward_impl::<CAPTURE>(scratch, layer_idx, state, trace_ptr);
4957 profiler.record_attention(layer_idx, attn_start.elapsed());
4958 } else {
4959 self.attention_forward_impl::<CAPTURE>(scratch, layer_idx, state, trace_ptr);
4960 }
4961
4962 kernel::add_avx(
4964 scratch.x.as_ptr(),
4965 scratch.att_out.as_ptr(),
4966 scratch.x.as_mut_ptr(),
4967 c,
4968 );
4969 if CAPTURE {
4970 scratch.train_trace_layers[layer_idx]
4971 .x_after_attn
4972 .copy_from(&scratch.x);
4973 }
4974
4975 kernel::layer_norm_avx(
4977 scratch.x.as_ptr(),
4978 self.blocks[layer_idx].ffn_norm_w.as_ptr(),
4979 self.blocks[layer_idx].ffn_norm_b.as_ptr(),
4980 scratch.x_normed.as_mut_ptr(),
4981 c,
4982 self.cfg.layer_norm_eps,
4983 );
4984 if CAPTURE {
4985 scratch.train_trace_layers[layer_idx]
4986 .ffn_norm
4987 .copy_from(&scratch.x_normed);
4988 }
4989
4990 if S::ENABLED {
4991 let ffn_start = Instant::now();
4992 self.ffn_forward_impl::<CAPTURE>(
4993 scratch,
4994 layer_idx,
4995 &mut state.layers[layer_idx],
4996 trace_ptr,
4997 );
4998 profiler.record_ffn(layer_idx, ffn_start.elapsed());
4999 } else {
5000 self.ffn_forward_impl::<CAPTURE>(
5001 scratch,
5002 layer_idx,
5003 &mut state.layers[layer_idx],
5004 trace_ptr,
5005 );
5006 }
5007
5008 kernel::add_avx(
5010 scratch.x.as_ptr(),
5011 scratch.ffn_out.as_ptr(),
5012 scratch.x.as_mut_ptr(),
5013 c,
5014 );
5015 if CAPTURE {
5016 scratch.train_trace_layers[layer_idx]
5017 .x_out
5018 .copy_from(&scratch.x);
5019 }
5020 }
5021
5022 kernel::layer_norm_avx(
5024 scratch.x.as_ptr(),
5025 self.ln_out_w.as_ptr(),
5026 self.ln_out_b.as_ptr(),
5027 scratch.x_normed.as_mut_ptr(),
5028 c,
5029 self.cfg.layer_norm_eps,
5030 );
5031
5032 kernel::gemv_avx(
5034 self.lm_head.as_ptr(),
5035 scratch.x_normed.as_ptr(),
5036 scratch.logits.as_mut_ptr(),
5037 self.cfg.vocab_size,
5038 c,
5039 );
5040 }
5041 if CAPTURE {
5042 scratch.train_v_first.copy_from(&state.v_first);
5043 }
5044
5045 scratch.logits.as_slice()
5046 }
5047
5048 #[inline(always)]
5049 unsafe fn attention_forward_impl<const CAPTURE: bool>(
5050 &self,
5051 scratch: &mut ScratchBuffers,
5052 layer_idx: usize,
5053 state: &mut State,
5054 trace: *mut LayerTrainTrace,
5055 ) {
5056 let attn = &self.blocks[layer_idx].attn;
5057 let layer_state = &mut state.layers[layer_idx];
5058 let c = self.cfg.hidden_size;
5059 let h = self.cfg.num_heads;
5060 let n = self.cfg.head_dim;
5061 let d_w = self.cfg.decay_low_rank;
5062 let d_a = self.cfg.a_low_rank;
5063 let d_g = self.cfg.g_low_rank;
5064 if CAPTURE {
5065 let tr = &mut *trace;
5066 tr.att_x_prev_old.copy_from(&layer_state.att_x_prev);
5067 tr.att_state_old.copy_from(&layer_state.att_state);
5068 }
5069
5070 kernel::token_shift_multi6_avx(
5071 scratch.x_normed.as_ptr(),
5072 layer_state.att_x_prev.as_ptr(),
5073 attn.x_r.as_ptr(),
5074 attn.x_w.as_ptr(),
5075 attn.x_k.as_ptr(),
5076 attn.x_v.as_ptr(),
5077 attn.x_a.as_ptr(),
5078 attn.x_g.as_ptr(),
5079 scratch.xr.as_mut_ptr(),
5080 scratch.xw.as_mut_ptr(),
5081 scratch.xk.as_mut_ptr(),
5082 scratch.xv.as_mut_ptr(),
5083 scratch.xa.as_mut_ptr(),
5084 scratch.xg.as_mut_ptr(),
5085 c,
5086 );
5087 if CAPTURE {
5088 let tr = &mut *trace;
5089 tr.xr.copy_from(&scratch.xr);
5090 tr.xw.copy_from(&scratch.xw);
5091 tr.xk.copy_from(&scratch.xk);
5092 tr.xv.copy_from(&scratch.xv);
5093 tr.xa.copy_from(&scratch.xa);
5094 tr.xg.copy_from(&scratch.xg);
5095 }
5096
5097 kernel::copy(
5099 scratch.x_normed.as_ptr(),
5100 layer_state.att_x_prev.as_mut_ptr(),
5101 c,
5102 );
5103
5104 let proj_size = c * c;
5107 kernel::gemv_avx(
5108 attn.rkv_proj.as_ptr(),
5109 scratch.xr.as_ptr(),
5110 scratch.r.as_mut_ptr(),
5111 c,
5112 c,
5113 );
5114 kernel::gemv_avx(
5115 attn.rkv_proj.as_ptr().add(proj_size),
5116 scratch.xk.as_ptr(),
5117 scratch.k.as_mut_ptr(),
5118 c,
5119 c,
5120 );
5121 kernel::gemv_avx(
5122 attn.rkv_proj.as_ptr().add(2 * proj_size),
5123 scratch.xv.as_ptr(),
5124 scratch.v.as_mut_ptr(),
5125 c,
5126 c,
5127 );
5128 if CAPTURE {
5129 let tr = &mut *trace;
5130 tr.r.copy_from(&scratch.r);
5131 tr.k_pre.copy_from(&scratch.k);
5132 tr.v_pre.copy_from(&scratch.v);
5133 }
5134
5135 kernel::gemv_avx(
5138 attn.w1.as_ptr(),
5139 scratch.xw.as_ptr(),
5140 scratch.w_lora_tmp.as_mut_ptr(),
5141 d_w,
5142 c,
5143 );
5144 kernel::tanh_avx(
5146 scratch.w_lora_tmp.as_ptr(),
5147 scratch.w_lora_tmp.as_mut_ptr(),
5148 d_w,
5149 );
5150 if CAPTURE {
5151 let tr = &mut *trace;
5152 tr.w_hidden.as_mut_slice()[0..d_w]
5153 .copy_from_slice(&scratch.w_lora_tmp.as_slice()[0..d_w]);
5154 }
5155 kernel::gemv_avx(
5157 attn.w2.as_ptr(),
5158 scratch.w_lora_tmp.as_ptr(),
5159 scratch.w_decay.as_mut_ptr(),
5160 c,
5161 d_w,
5162 );
5163 kernel::add_avx(
5165 scratch.w_decay.as_ptr(),
5166 attn.w0.as_ptr(),
5167 scratch.w_decay.as_mut_ptr(),
5168 c,
5169 );
5170 if CAPTURE {
5171 let tr = &mut *trace;
5172 tr.w_pre.copy_from(&scratch.w_decay);
5173 }
5174 let inv_sqrt_e = 1.0 / std::f32::consts::E.sqrt();
5176 kernel::sigmoid_exp_neg_scaled_avx(
5177 scratch.w_decay.as_ptr(),
5178 scratch.w_decay.as_mut_ptr(),
5179 if CAPTURE {
5180 (*trace).w_sigmoid.as_mut_ptr()
5181 } else {
5182 std::ptr::null_mut()
5183 },
5184 inv_sqrt_e,
5185 c,
5186 );
5187 if CAPTURE {
5188 let tr = &mut *trace;
5189 tr.w_decay.copy_from(&scratch.w_decay);
5190 }
5191
5192 kernel::gemv_avx(
5194 attn.a1.as_ptr(),
5195 scratch.xa.as_ptr(),
5196 scratch.w_lora_tmp.as_mut_ptr(),
5197 d_a,
5198 c,
5199 );
5200 if CAPTURE {
5201 let tr = &mut *trace;
5202 tr.a_hidden.as_mut_slice()[0..d_a]
5203 .copy_from_slice(&scratch.w_lora_tmp.as_slice()[0..d_a]);
5204 }
5205 kernel::gemv_avx(
5206 attn.a2.as_ptr(),
5207 scratch.w_lora_tmp.as_ptr(),
5208 scratch.a.as_mut_ptr(),
5209 c,
5210 d_a,
5211 );
5212 kernel::add_avx(
5213 scratch.a.as_ptr(),
5214 attn.a0.as_ptr(),
5215 scratch.a.as_mut_ptr(),
5216 c,
5217 );
5218 kernel::sigmoid_avx(scratch.a.as_ptr(), scratch.a.as_mut_ptr(), c);
5219 if CAPTURE {
5220 let tr = &mut *trace;
5221 tr.a.copy_from(&scratch.a);
5222 }
5223
5224 kernel::gemv_avx(
5226 attn.g1.as_ptr(),
5227 scratch.xg.as_ptr(),
5228 scratch.w_lora_tmp.as_mut_ptr(),
5229 d_g,
5230 c,
5231 );
5232 kernel::sigmoid_avx(
5233 scratch.w_lora_tmp.as_ptr(),
5234 scratch.w_lora_tmp.as_mut_ptr(),
5235 d_g,
5236 );
5237 if CAPTURE {
5238 let tr = &mut *trace;
5239 tr.g_hidden.as_mut_slice()[0..d_g]
5240 .copy_from_slice(&scratch.w_lora_tmp.as_slice()[0..d_g]);
5241 }
5242 kernel::gemv_avx(
5243 attn.g2.as_ptr(),
5244 scratch.w_lora_tmp.as_ptr(),
5245 scratch.g.as_mut_ptr(),
5246 c,
5247 d_g,
5248 );
5249 if CAPTURE {
5250 let tr = &mut *trace;
5251 tr.g.copy_from(&scratch.g);
5252 }
5253
5254 if layer_idx == 0 {
5256 state.v_first.copy_from(&scratch.v);
5258 state.v_first_set = true;
5259 if CAPTURE {
5260 let tr = &mut *trace;
5261 tr.uses_v_residual = false;
5262 tr.v.copy_from(&scratch.v);
5263 }
5264 } else if state.v_first_set
5265 && let (Some(v1), Some(v2), Some(v0)) = (&attn.v1, &attn.v2, &attn.v0)
5266 {
5267 let d_v = self.cfg.v_low_rank;
5268 kernel::gemv_avx(
5270 v1.as_ptr(),
5271 scratch.xv.as_ptr(),
5272 scratch.w_lora_tmp.as_mut_ptr(),
5273 d_v,
5274 c,
5275 );
5276 if CAPTURE {
5277 let tr = &mut *trace;
5278 tr.v_hidden.as_mut_slice()[0..d_v]
5279 .copy_from_slice(&scratch.w_lora_tmp.as_slice()[0..d_v]);
5280 }
5281 kernel::gemv_avx(
5282 v2.as_ptr(),
5283 scratch.w_lora_tmp.as_ptr(),
5284 scratch.att_out.as_mut_ptr(), c,
5286 d_v,
5287 );
5288 kernel::add_avx(
5289 scratch.att_out.as_ptr(),
5290 v0.as_ptr(),
5291 scratch.att_out.as_mut_ptr(),
5292 c,
5293 );
5294 kernel::sigmoid_avx(scratch.att_out.as_ptr(), scratch.att_out.as_mut_ptr(), c);
5295 if CAPTURE {
5296 let tr = &mut *trace;
5297 tr.uses_v_residual = true;
5298 tr.nu.copy_from(&scratch.att_out);
5299 }
5300 for i in 0..c {
5302 let nu = scratch.att_out[i];
5303 scratch.v[i] += (state.v_first[i] - scratch.v[i]) * nu;
5304 }
5305 if CAPTURE {
5306 let tr = &mut *trace;
5307 tr.v.copy_from(&scratch.v);
5308 }
5309 } else if CAPTURE {
5310 let tr = &mut *trace;
5311 tr.uses_v_residual = false;
5312 tr.v.copy_from(&scratch.v);
5313 }
5314
5315 kernel::mul_avx(
5317 scratch.k.as_ptr(),
5318 attn.k_k.as_ptr(),
5319 scratch.kk.as_mut_ptr(),
5320 c,
5321 );
5322 if CAPTURE {
5323 let tr = &mut *trace;
5324 tr.kk_pre.copy_from(&scratch.kk);
5325 }
5326 for head in 0..h {
5328 let offset = head * n;
5329 kernel::l2_normalize_avx(
5330 scratch.kk.as_ptr().add(offset),
5331 scratch.kk.as_mut_ptr().add(offset),
5332 n,
5333 1e-12,
5334 );
5335 }
5336 if CAPTURE {
5337 let tr = &mut *trace;
5338 tr.kk.copy_from(&scratch.kk);
5339 }
5340
5341 for i in 0..c {
5343 let scale = 1.0 + (scratch.a[i] - 1.0) * attn.k_a[i];
5344 scratch.k[i] *= scale;
5345 }
5346 if CAPTURE {
5347 let tr = &mut *trace;
5348 tr.k.copy_from(&scratch.k);
5349 }
5350
5351 kernel::rwkv7_wkv_update_avx(
5353 layer_state.att_state.as_mut_ptr(),
5354 scratch.w_decay.as_ptr(),
5355 scratch.k.as_ptr(),
5356 scratch.v.as_ptr(),
5357 scratch.kk.as_ptr(),
5358 scratch.a.as_ptr(),
5359 scratch.r.as_ptr(),
5360 scratch.y.as_mut_ptr(),
5361 h,
5362 n,
5363 );
5364 if CAPTURE {
5365 let tr = &mut *trace;
5366 tr.y_wkv.copy_from(&scratch.y);
5367 }
5368
5369 kernel::group_norm_avx(
5371 scratch.y.as_ptr(),
5372 attn.g_norm_w.as_ptr(),
5373 attn.g_norm_b.as_ptr(),
5374 scratch.y.as_mut_ptr(),
5375 h,
5376 n,
5377 self.cfg.group_norm_eps,
5378 );
5379 if CAPTURE {
5380 let tr = &mut *trace;
5381 tr.y_gn.copy_from(&scratch.y);
5382 }
5383
5384 for head in 0..h {
5386 let offset = head * n;
5387 let mut alpha = 0.0f32;
5388 for j in 0..n {
5389 alpha += scratch.r[offset + j] * scratch.k[offset + j] * attn.r_k[head * n + j];
5390 }
5391 if CAPTURE {
5392 let tr = &mut *trace;
5393 tr.alpha[head] = alpha;
5394 }
5395 for j in 0..n {
5396 scratch.y[offset + j] += alpha * scratch.v[offset + j];
5397 }
5398 }
5399 if CAPTURE {
5400 let tr = &mut *trace;
5401 tr.y_head.copy_from(&scratch.y);
5402 }
5403
5404 kernel::mul_avx(
5406 scratch.y.as_ptr(),
5407 scratch.g.as_ptr(),
5408 scratch.y.as_mut_ptr(),
5409 c,
5410 );
5411 if CAPTURE {
5412 let tr = &mut *trace;
5413 tr.y_gate.copy_from(&scratch.y);
5414 }
5415
5416 kernel::gemv_avx(
5418 attn.o_proj.as_ptr(),
5419 scratch.y.as_ptr(),
5420 scratch.att_out.as_mut_ptr(),
5421 c,
5422 c,
5423 );
5424 if CAPTURE {
5425 let tr = &mut *trace;
5426 tr.att_out.copy_from(&scratch.att_out);
5427 }
5428 }
5429
5430 #[inline(always)]
5431 unsafe fn ffn_forward_impl<const CAPTURE: bool>(
5432 &self,
5433 scratch: &mut ScratchBuffers,
5434 layer_idx: usize,
5435 layer_state: &mut LayerState,
5436 trace: *mut LayerTrainTrace,
5437 ) {
5438 let ffn = &self.blocks[layer_idx].ffn;
5439 let c = self.cfg.hidden_size;
5440 let i = self.cfg.intermediate_size;
5441 if CAPTURE {
5442 let tr = &mut *trace;
5443 tr.ffn_x_prev_old.copy_from(&layer_state.ffn_x_prev);
5444 }
5445
5446 kernel::token_shift_avx(
5448 scratch.x_normed.as_ptr(),
5449 layer_state.ffn_x_prev.as_ptr(),
5450 ffn.x_k.as_ptr(),
5451 scratch.xk.as_mut_ptr(),
5452 c,
5453 );
5454 if CAPTURE {
5455 let tr = &mut *trace;
5456 tr.ffn_xk.copy_from(&scratch.xk);
5457 }
5458
5459 kernel::copy(
5461 scratch.x_normed.as_ptr(),
5462 layer_state.ffn_x_prev.as_mut_ptr(),
5463 c,
5464 );
5465
5466 kernel::gemv_avx(
5468 ffn.key_w.as_ptr(),
5469 scratch.xk.as_ptr(),
5470 scratch.ffn_k.as_mut_ptr(),
5471 i,
5472 c,
5473 );
5474 if CAPTURE {
5475 let tr = &mut *trace;
5476 tr.ffn_pre.copy_from(&scratch.ffn_k);
5477 }
5478 kernel::relu_squared_avx(scratch.ffn_k.as_ptr(), scratch.ffn_k.as_mut_ptr(), i);
5479 if CAPTURE {
5480 let tr = &mut *trace;
5481 tr.ffn_k.copy_from(&scratch.ffn_k);
5482 }
5483
5484 kernel::gemv_avx(
5486 ffn.value_w.as_ptr(),
5487 scratch.ffn_k.as_ptr(),
5488 scratch.ffn_out.as_mut_ptr(),
5489 c,
5490 i,
5491 );
5492 if CAPTURE {
5493 let tr = &mut *trace;
5494 tr.ffn_out.copy_from(&scratch.ffn_out);
5495 }
5496 }
5497}
5498
5499#[allow(clippy::needless_range_loop)]
5500fn layer_norm_backward(
5501 input: &[f32],
5502 weight: &[f32],
5503 grad_out: &[f32],
5504 eps: f32,
5505 grad_input: &mut [f32],
5506 grad_weight: &mut [f32],
5507 grad_bias: &mut [f32],
5508) {
5509 let n = input
5510 .len()
5511 .min(weight.len())
5512 .min(grad_out.len())
5513 .min(grad_input.len())
5514 .min(grad_weight.len())
5515 .min(grad_bias.len());
5516 if n == 0 {
5517 return;
5518 }
5519 let nf = n as f32;
5520 let mut mean = 0.0f32;
5521 for &x in &input[0..n] {
5522 mean += x;
5523 }
5524 mean /= nf;
5525 let mut var = 0.0f32;
5526 for &x in &input[0..n] {
5527 let d = x - mean;
5528 var += d * d;
5529 }
5530 var /= nf;
5531 let inv_std = (var + eps).sqrt().recip();
5532 let mut sum_gw = 0.0f32;
5533 let mut sum_gw_xhat = 0.0f32;
5534 for i in 0..n {
5535 let xhat = (input[i] - mean) * inv_std;
5536 let gw = grad_out[i] * weight[i];
5537 grad_weight[i] = grad_out[i] * xhat;
5538 grad_bias[i] = grad_out[i];
5539 sum_gw += gw;
5540 sum_gw_xhat += gw * xhat;
5541 }
5542 for i in 0..n {
5543 let xhat = (input[i] - mean) * inv_std;
5544 let gw = grad_out[i] * weight[i];
5545 grad_input[i] = (gw * nf - sum_gw - xhat * sum_gw_xhat) * inv_std / nf;
5546 }
5547}
5548
5549#[allow(clippy::needless_range_loop, clippy::too_many_arguments)]
5550fn group_norm_backward(
5551 input: &[f32],
5552 weight: &[f32],
5553 grad_out: &[f32],
5554 num_groups: usize,
5555 group_size: usize,
5556 eps: f32,
5557 grad_input: &mut [f32],
5558 grad_weight: &mut [f32],
5559 grad_bias: &mut [f32],
5560) {
5561 let c = input
5562 .len()
5563 .min(weight.len())
5564 .min(grad_out.len())
5565 .min(grad_input.len())
5566 .min(grad_weight.len())
5567 .min(grad_bias.len());
5568 if c == 0 || num_groups == 0 || group_size == 0 {
5569 return;
5570 }
5571 grad_input[0..c].fill(0.0);
5572 grad_weight[0..c].fill(0.0);
5573 grad_bias[0..c].fill(0.0);
5574 let g = num_groups.min(c / group_size);
5575 let n = group_size as f32;
5576 for group in 0..g {
5577 let off = group * group_size;
5578 let end = (off + group_size).min(c);
5579 let len = end - off;
5580 if len == 0 {
5581 continue;
5582 }
5583 let mut mean = 0.0f32;
5584 for idx in off..end {
5585 mean += input[idx];
5586 }
5587 mean /= len as f32;
5588 let mut var = 0.0f32;
5589 for idx in off..end {
5590 let d = input[idx] - mean;
5591 var += d * d;
5592 }
5593 var /= len as f32;
5594 let inv_std = (var + eps).sqrt().recip();
5595 let mut sum_gw = 0.0f32;
5596 let mut sum_gw_xhat = 0.0f32;
5597 for idx in off..end {
5598 let xhat = (input[idx] - mean) * inv_std;
5599 let gw = grad_out[idx] * weight[idx];
5600 grad_weight[idx] += grad_out[idx] * xhat;
5601 grad_bias[idx] += grad_out[idx];
5602 sum_gw += gw;
5603 sum_gw_xhat += gw * xhat;
5604 }
5605 for idx in off..end {
5606 let xhat = (input[idx] - mean) * inv_std;
5607 let gw = grad_out[idx] * weight[idx];
5608 grad_input[idx] = (gw * n - sum_gw - xhat * sum_gw_xhat) * inv_std / n;
5609 }
5610 }
5611}
5612
5613fn l2_normalize_backward(
5614 x: &[f32],
5615 y: &[f32],
5616 grad_out: &[f32],
5617 min_norm: f32,
5618 grad_input: &mut [f32],
5619) {
5620 let n = x
5621 .len()
5622 .min(y.len())
5623 .min(grad_out.len())
5624 .min(grad_input.len());
5625 if n == 0 {
5626 return;
5627 }
5628 let mut norm_sq = 0.0f32;
5629 for &v in &x[0..n] {
5630 norm_sq += v * v;
5631 }
5632 let norm_raw = norm_sq.sqrt();
5633 if norm_raw <= min_norm {
5634 let inv = min_norm.recip();
5635 for i in 0..n {
5636 grad_input[i] = grad_out[i] * inv;
5637 }
5638 return;
5639 }
5640 let norm = norm_raw;
5641 let mut dot = 0.0f32;
5642 for i in 0..n {
5643 dot += grad_out[i] * y[i];
5644 }
5645 let inv = norm.recip();
5646 for i in 0..n {
5647 grad_input[i] = (grad_out[i] - y[i] * dot) * inv;
5648 }
5649}
5650
5651#[inline(always)]
5652fn add_vec_grad(dst: &mut [f32], src: &[f32]) {
5653 let n = dst.len().min(src.len());
5654 for i in 0..n {
5655 dst[i] += src[i];
5656 }
5657}
5658
5659#[inline(always)]
5660#[allow(clippy::needless_range_loop)]
5661fn add_outer_grad(dst: &mut [f32], rows: usize, cols: usize, left: &[f32], right: &[f32]) {
5662 let rows = rows.min(left.len());
5663 let cols = cols.min(right.len());
5664 let n = dst.len();
5665 if rows == 0 || cols == 0 || n == 0 {
5666 return;
5667 }
5668 for r in 0..rows {
5669 let g = left[r];
5670 if g == 0.0 {
5671 continue;
5672 }
5673 let off = r * cols;
5674 if off >= n {
5675 break;
5676 }
5677 let row_cols = cols.min(n - off);
5678 for c in 0..row_cols {
5679 dst[off + c] += g * right[c];
5680 }
5681 }
5682}
5683
5684#[inline(always)]
5685fn sgd_vec_update(param: &mut [f32], grad: &[f32], lr: f32, clip: f32) {
5686 let n = param.len().min(grad.len());
5687 if n == 0 {
5688 return;
5689 }
5690 if clip > 0.0 {
5691 for i in 0..n {
5692 param[i] += lr * grad[i].clamp(-clip, clip);
5693 }
5694 } else {
5695 for i in 0..n {
5696 param[i] += lr * grad[i];
5697 }
5698 }
5699}
5700
5701#[inline(always)]
5702#[allow(clippy::needless_range_loop)]
5703fn sgd_outer_update(
5704 param: &mut [f32],
5705 rows: usize,
5706 cols: usize,
5707 left: &[f32],
5708 right: &[f32],
5709 lr: f32,
5710 clip: f32,
5711) {
5712 let rows = rows.min(left.len());
5713 let cols = cols.min(right.len());
5714 let n = param.len();
5715 if rows == 0 || cols == 0 || n == 0 {
5716 return;
5717 }
5718 for r in 0..rows {
5719 let g = left[r];
5720 let off = r * cols;
5721 if off >= n {
5722 break;
5723 }
5724 let row_cols = cols.min(n - off);
5725 if clip > 0.0 {
5726 for c in 0..row_cols {
5727 param[off + c] += lr * (g * right[c]).clamp(-clip, clip);
5728 }
5729 } else {
5730 for c in 0..row_cols {
5731 param[off + c] += lr * g * right[c];
5732 }
5733 }
5734 }
5735}
5736
5737#[inline(always)]
5738#[allow(clippy::needless_range_loop, clippy::too_many_arguments)]
5739fn fused_sgd_head_backward_update(
5740 param: &mut [f32],
5741 rows: usize,
5742 cols: usize,
5743 left: &[f32],
5744 right: &[f32],
5745 grad_input: &mut [f32],
5746 lr: f32,
5747 clip: f32,
5748) {
5749 let rows = rows.min(left.len());
5750 let cols = cols.min(right.len()).min(grad_input.len());
5751 let n = param.len();
5752 if rows == 0 || cols == 0 || n == 0 {
5753 return;
5754 }
5755 let do_clip = clip > 0.0;
5756 let lr8 = f32x8::splat(lr);
5757 for row in 0..rows {
5758 let g = left[row];
5759 if g == 0.0 {
5760 continue;
5761 }
5762 let off = row * cols;
5763 if off >= n {
5764 break;
5765 }
5766 let row_cols = cols.min(n - off);
5767 if do_clip {
5768 for col in 0..row_cols {
5769 let idx = off + col;
5770 let w_old = param[idx];
5771 grad_input[col] += w_old * g;
5772 param[idx] = w_old + lr * (g * right[col]).clamp(-clip, clip);
5773 }
5774 continue;
5775 }
5776 let mut col = 0usize;
5777 unsafe {
5778 let g8 = f32x8::splat(g);
5779 while col + 8 <= row_cols {
5780 let idx = off + col;
5781 let wv = param.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5782 let rv = right.as_ptr().add(col).cast::<f32x8>().read_unaligned();
5783 let giv = grad_input
5784 .as_ptr()
5785 .add(col)
5786 .cast::<f32x8>()
5787 .read_unaligned();
5788 grad_input
5789 .as_mut_ptr()
5790 .add(col)
5791 .cast::<f32x8>()
5792 .write_unaligned(giv + wv * g8);
5793 param
5794 .as_mut_ptr()
5795 .add(idx)
5796 .cast::<f32x8>()
5797 .write_unaligned(wv + (g8 * rv) * lr8);
5798 col += 8;
5799 }
5800 }
5801 while col < row_cols {
5802 let idx = off + col;
5803 let w_old = param[idx];
5804 grad_input[col] += w_old * g;
5805 param[idx] = w_old + lr * g * right[col];
5806 col += 1;
5807 }
5808 }
5809}
5810
5811#[inline(always)]
5812fn apply_adam_vec_update(
5813 param: &mut [f32],
5814 grad: &[f32],
5815 adam: &mut AdamTensorState,
5816 step: &AdamStep,
5817) {
5818 let n = param
5819 .len()
5820 .min(grad.len())
5821 .min(adam.m.len())
5822 .min(adam.v.len());
5823 if n == 0 {
5824 return;
5825 }
5826 apply_adam_vec_update_raw(
5827 &mut param[0..n],
5828 &grad[0..n],
5829 &mut adam.m.as_mut_slice()[0..n],
5830 &mut adam.v.as_mut_slice()[0..n],
5831 step,
5832 );
5833}
5834
5835#[inline(always)]
5836fn apply_adam_vec_update_raw(
5837 param: &mut [f32],
5838 grad: &[f32],
5839 m: &mut [f32],
5840 v: &mut [f32],
5841 step: &AdamStep,
5842) {
5843 let n = param.len().min(grad.len()).min(m.len()).min(v.len());
5844 if n == 0 {
5845 return;
5846 }
5847 let b1 = step.b1;
5848 let b2 = step.b2;
5849 let one_b1 = 1.0 - b1;
5850 let one_b2 = 1.0 - b2;
5851 let inv_bc1 = 1.0 / step.bias_corr1;
5852 let inv_bc2 = 1.0 / step.bias_corr2;
5853 let do_clip = step.clip > 0.0;
5854 let clip = step.clip;
5855 if do_clip {
5856 for idx in 0..n {
5857 let g = grad[idx].clamp(-clip, clip);
5858 let mm = b1 * m[idx] + one_b1 * g;
5859 let vv = b2 * v[idx] + one_b2 * g * g;
5860 m[idx] = mm;
5861 v[idx] = vv;
5862 let m_hat = mm * inv_bc1;
5863 let v_hat = vv * inv_bc2;
5864 param[idx] += step.lr * m_hat / (v_hat.sqrt() + step.eps);
5865 }
5866 return;
5867 }
5868 let mut idx = 0usize;
5869 unsafe {
5870 let b1v = f32x8::splat(b1);
5871 let b2v = f32x8::splat(b2);
5872 let one_b1v = f32x8::splat(one_b1);
5873 let one_b2v = f32x8::splat(one_b2);
5874 let inv_bc1v = f32x8::splat(inv_bc1);
5875 let inv_bc2v = f32x8::splat(inv_bc2);
5876 let lrv = f32x8::splat(step.lr);
5877 let epsv = f32x8::splat(step.eps);
5878 while idx + 8 <= n {
5879 let gv = grad.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5880 let mv = m.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5881 let vv = v.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5882 let mm = mv * b1v + gv * one_b1v;
5883 let vv2 = vv * b2v + (gv * gv) * one_b2v;
5884 m.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(mm);
5885 v.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(vv2);
5886 let pv = param.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5887 let upd = ((mm * inv_bc1v) / ((vv2 * inv_bc2v).sqrt() + epsv)) * lrv;
5888 param
5889 .as_mut_ptr()
5890 .add(idx)
5891 .cast::<f32x8>()
5892 .write_unaligned(pv + upd);
5893 idx += 8;
5894 }
5895 }
5896 while idx < n {
5897 let g = grad[idx];
5898 let mm = b1 * m[idx] + one_b1 * g;
5899 let vv = b2 * v[idx] + one_b2 * g * g;
5900 m[idx] = mm;
5901 v[idx] = vv;
5902 let m_hat = mm * inv_bc1;
5903 let v_hat = vv * inv_bc2;
5904 param[idx] += step.lr * m_hat / (v_hat.sqrt() + step.eps);
5905 idx += 1;
5906 }
5907}
5908
5909#[inline(always)]
5910#[allow(clippy::needless_range_loop, clippy::too_many_arguments)]
5911fn fused_adam_head_backward_update(
5912 param: &mut [f32],
5913 rows: usize,
5914 cols: usize,
5915 left: &[f32],
5916 right: &[f32],
5917 grad_input: &mut [f32],
5918 m: &mut [f32],
5919 v: &mut [f32],
5920 step: &AdamStep,
5921) {
5922 let rows = rows.min(left.len());
5923 let cols = cols.min(right.len()).min(grad_input.len());
5924 let n = param.len().min(m.len()).min(v.len());
5925 if rows == 0 || cols == 0 || n == 0 {
5926 return;
5927 }
5928 let b1 = step.b1;
5929 let b2 = step.b2;
5930 let one_b1 = 1.0 - b1;
5931 let one_b2 = 1.0 - b2;
5932 let inv_bc1 = 1.0 / step.bias_corr1;
5933 let inv_bc2 = 1.0 / step.bias_corr2;
5934 let do_clip = step.clip > 0.0;
5935 let clip = step.clip;
5936 let b1v = f32x8::splat(b1);
5937 let b2v = f32x8::splat(b2);
5938 let one_b1v = f32x8::splat(one_b1);
5939 let one_b2v = f32x8::splat(one_b2);
5940 let inv_bc1v = f32x8::splat(inv_bc1);
5941 let inv_bc2v = f32x8::splat(inv_bc2);
5942 let epsv = f32x8::splat(step.eps);
5943 let lrv = f32x8::splat(step.lr);
5944 for row in 0..rows {
5945 let g = left[row];
5946 if g == 0.0 {
5947 continue;
5948 }
5949 let off = row * cols;
5950 if off >= n {
5951 break;
5952 }
5953 let row_cols = cols.min(n - off);
5954 if do_clip {
5955 for col in 0..row_cols {
5956 let idx = off + col;
5957 let w_old = param[idx];
5958 grad_input[col] += w_old * g;
5959 let gg = (g * right[col]).clamp(-clip, clip);
5960 let mm = b1 * m[idx] + one_b1 * gg;
5961 let vv = b2 * v[idx] + one_b2 * gg * gg;
5962 m[idx] = mm;
5963 v[idx] = vv;
5964 let m_hat = mm * inv_bc1;
5965 let v_hat = vv * inv_bc2;
5966 param[idx] = w_old + step.lr * m_hat / (v_hat.sqrt() + step.eps);
5967 }
5968 continue;
5969 }
5970 let mut col = 0usize;
5971 unsafe {
5972 let g8 = f32x8::splat(g);
5973 while col + 8 <= row_cols {
5974 let idx = off + col;
5975 let wv = param.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5976 let rv = right.as_ptr().add(col).cast::<f32x8>().read_unaligned();
5977 let giv = grad_input
5978 .as_ptr()
5979 .add(col)
5980 .cast::<f32x8>()
5981 .read_unaligned();
5982 grad_input
5983 .as_mut_ptr()
5984 .add(col)
5985 .cast::<f32x8>()
5986 .write_unaligned(giv + wv * g8);
5987 let gv = g8 * rv;
5988 let mv = m.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5989 let vv = v.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
5990 let mm = mv * b1v + gv * one_b1v;
5991 let vv2 = vv * b2v + (gv * gv) * one_b2v;
5992 m.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(mm);
5993 v.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(vv2);
5994 let upd = ((mm * inv_bc1v) / ((vv2 * inv_bc2v).sqrt() + epsv)) * lrv;
5995 param
5996 .as_mut_ptr()
5997 .add(idx)
5998 .cast::<f32x8>()
5999 .write_unaligned(wv + upd);
6000 col += 8;
6001 }
6002 }
6003 while col < row_cols {
6004 let idx = off + col;
6005 let w_old = param[idx];
6006 grad_input[col] += w_old * g;
6007 let gg = g * right[col];
6008 let mm = b1 * m[idx] + one_b1 * gg;
6009 let vv = b2 * v[idx] + one_b2 * gg * gg;
6010 m[idx] = mm;
6011 v[idx] = vv;
6012 let m_hat = mm * inv_bc1;
6013 let v_hat = vv * inv_bc2;
6014 param[idx] = w_old + step.lr * m_hat / (v_hat.sqrt() + step.eps);
6015 col += 1;
6016 }
6017 }
6018}
6019
6020#[inline(always)]
6021fn apply_adam_outer_update(
6022 param: &mut [f32],
6023 rows: usize,
6024 cols: usize,
6025 left: &[f32],
6026 right: &[f32],
6027 adam: &mut AdamTensorState,
6028 step: &AdamStep,
6029) {
6030 let n = param.len().min(adam.m.len()).min(adam.v.len());
6031 if n == 0 {
6032 return;
6033 }
6034 apply_adam_outer_update_raw(
6035 &mut param[0..n],
6036 rows,
6037 cols,
6038 left,
6039 right,
6040 &mut adam.m.as_mut_slice()[0..n],
6041 &mut adam.v.as_mut_slice()[0..n],
6042 step,
6043 );
6044}
6045
6046#[allow(clippy::too_many_arguments)]
6047#[inline(always)]
6048#[allow(clippy::needless_range_loop)]
6049fn apply_adam_outer_update_raw(
6050 param: &mut [f32],
6051 rows: usize,
6052 cols: usize,
6053 left: &[f32],
6054 right: &[f32],
6055 m: &mut [f32],
6056 v: &mut [f32],
6057 step: &AdamStep,
6058) {
6059 let rows = rows.min(left.len());
6060 let cols = cols.min(right.len());
6061 let n = param.len().min(m.len()).min(v.len());
6062 if rows == 0 || cols == 0 || n == 0 {
6063 return;
6064 }
6065 let b1 = step.b1;
6066 let b2 = step.b2;
6067 let one_b1 = 1.0 - b1;
6068 let one_b2 = 1.0 - b2;
6069 let inv_bc1 = 1.0 / step.bias_corr1;
6070 let inv_bc2 = 1.0 / step.bias_corr2;
6071 let do_clip = step.clip > 0.0;
6072 let clip = step.clip;
6073 let b1v = f32x8::splat(b1);
6074 let b2v = f32x8::splat(b2);
6075 let one_b1v = f32x8::splat(one_b1);
6076 let one_b2v = f32x8::splat(one_b2);
6077 let inv_bc1v = f32x8::splat(inv_bc1);
6078 let inv_bc2v = f32x8::splat(inv_bc2);
6079 let epsv = f32x8::splat(step.eps);
6080 let lrv = f32x8::splat(step.lr);
6081 for row in 0..rows {
6082 let g_row = left[row];
6083 let off = row * cols;
6084 if off >= n {
6085 break;
6086 }
6087 let row_cols = (n - off).min(cols);
6088 if do_clip {
6089 for col in 0..row_cols {
6090 let idx = off + col;
6091 let g = (g_row * right[col]).clamp(-clip, clip);
6092 let mm = b1 * m[idx] + one_b1 * g;
6093 let vv = b2 * v[idx] + one_b2 * g * g;
6094 m[idx] = mm;
6095 v[idx] = vv;
6096 let m_hat = mm * inv_bc1;
6097 let v_hat = vv * inv_bc2;
6098 param[idx] += step.lr * m_hat / (v_hat.sqrt() + step.eps);
6099 }
6100 continue;
6101 }
6102 let mut col = 0usize;
6103 unsafe {
6104 let g8 = f32x8::splat(g_row);
6105 while col + 8 <= row_cols {
6106 let idx = off + col;
6107 let rv = right.as_ptr().add(col).cast::<f32x8>().read_unaligned();
6108 let gv = g8 * rv;
6109 let mv = m.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
6110 let vv = v.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
6111 let mm = mv * b1v + gv * one_b1v;
6112 let vv2 = vv * b2v + (gv * gv) * one_b2v;
6113 m.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(mm);
6114 v.as_mut_ptr().add(idx).cast::<f32x8>().write_unaligned(vv2);
6115 let pv = param.as_ptr().add(idx).cast::<f32x8>().read_unaligned();
6116 let upd = ((mm * inv_bc1v) / ((vv2 * inv_bc2v).sqrt() + epsv)) * lrv;
6117 param
6118 .as_mut_ptr()
6119 .add(idx)
6120 .cast::<f32x8>()
6121 .write_unaligned(pv + upd);
6122 col += 8;
6123 }
6124 }
6125 while col < row_cols {
6126 let idx = off + col;
6127 let g = g_row * right[col];
6128 let mm = b1 * m[idx] + one_b1 * g;
6129 let vv = b2 * v[idx] + one_b2 * g * g;
6130 m[idx] = mm;
6131 v[idx] = vv;
6132 let m_hat = mm * inv_bc1;
6133 let v_hat = vv * inv_bc2;
6134 param[idx] += step.lr * m_hat / (v_hat.sqrt() + step.eps);
6135 col += 1;
6136 }
6137 }
6138}
6139
6140struct RwkvRng {
6141 state: u64,
6142}
6143
6144impl RwkvRng {
6145 fn new(seed: u64) -> Self {
6146 Self {
6147 state: seed ^ 0x9E37_79B9_7F4A_7C15,
6148 }
6149 }
6150
6151 #[inline]
6152 fn next_u32(&mut self) -> u32 {
6153 self.state = self
6154 .state
6155 .wrapping_mul(6_364_136_223_846_793_005)
6156 .wrapping_add(1);
6157 (self.state >> 32) as u32
6158 }
6159
6160 #[inline]
6161 fn next_f32(&mut self) -> f32 {
6162 let v = self.next_u32() as f32;
6163 v * (1.0 / (u32::MAX as f32))
6164 }
6165}
6166
6167#[inline]
6168fn init_uniform(t: &mut Tensor1D, rng: &mut RwkvRng, scale: f32) {
6169 let s = t.as_mut_slice();
6170 for v in s {
6171 let r = rng.next_f32() - 0.5;
6172 *v = r * 2.0 * scale;
6173 }
6174}
6175
6176#[inline]
6177fn init_centered(t: &mut Tensor1D, rng: &mut RwkvRng, center: f32, scale: f32) {
6178 let s = t.as_mut_slice();
6179 for v in s {
6180 let r = rng.next_f32() - 0.5;
6181 *v = center + r * 2.0 * scale;
6182 }
6183}
6184
6185#[inline]
6186fn init_const(t: &mut Tensor1D, value: f32) {
6187 t.as_mut_slice().fill(value);
6188}
6189
6190#[cfg(test)]
6191mod tests {
6192 use super::*;
6193 use std::path::PathBuf;
6194
6195 fn test_cfg() -> Config {
6196 Config {
6197 vocab_size: 256,
6198 hidden_size: 64,
6199 num_layers: 1,
6200 num_heads: 1,
6201 head_dim: 64,
6202 intermediate_size: 64,
6203 layer_norm_eps: 1e-5,
6204 group_norm_eps: 64e-5,
6205 decay_low_rank: 8,
6206 a_low_rank: 8,
6207 v_low_rank: 8,
6208 g_low_rank: 8,
6209 }
6210 }
6211
6212 fn temp_path(prefix: &str, ext: &str) -> PathBuf {
6213 let now = std::time::SystemTime::now()
6214 .duration_since(std::time::UNIX_EPOCH)
6215 .unwrap_or_default()
6216 .as_nanos();
6217 std::env::temp_dir().join(format!("{prefix}_{}_{}.{}", std::process::id(), now, ext))
6218 }
6219
6220 fn softmax_loss(logits: &[f32], target: u8) -> f64 {
6221 let max_logit = logits
6222 .iter()
6223 .copied()
6224 .fold(f32::NEG_INFINITY, |a, b| a.max(b));
6225 let mut sum = 0.0f64;
6226 for &z in logits {
6227 sum += ((z - max_logit) as f64).exp();
6228 }
6229 let p = ((logits[target as usize] - max_logit) as f64).exp() / sum.max(1e-300);
6230 -p.max(1e-300).ln()
6231 }
6232
6233 fn segment_loss(model: &Model, cfg: &Config, steps: &[(u32, u8)]) -> f64 {
6234 if steps.is_empty() {
6235 return 0.0;
6236 }
6237 let mut scratch = ScratchBuffers::new(cfg);
6238 let mut state = model.new_state();
6239 let mut loss = 0.0f64;
6240 for &(input, target) in steps {
6241 let logits = model.forward(&mut scratch, input, &mut state);
6242 loss += softmax_loss(logits, target);
6243 }
6244 loss / (steps.len() as f64)
6245 }
6246
6247 fn segment_grads(model: &Model, cfg: &Config, steps: &[(u32, u8)]) -> FullGradState {
6248 let mut scratch = ScratchBuffers::new(cfg);
6249 let mut state = model.new_state();
6250 let mut states = Vec::with_capacity(steps.len() + 1);
6251 let mut traces = Vec::with_capacity(steps.len());
6252 let mut pdfs = Vec::with_capacity(steps.len());
6253 states.push(state.clone());
6254 for &(input, _) in steps {
6255 scratch.set_capture_train_trace(true);
6256 let logits = model.forward(&mut scratch, input, &mut state);
6257 let mut pdf = vec![0.0f64; cfg.vocab_size];
6258 super::super::super::softmax_pdf_floor_with_bias(logits, None, &mut pdf);
6259 pdfs.push(pdf);
6260 traces.push(TokenTrainTrace::from_scratch(&scratch));
6261 states.push(state.clone());
6262 }
6263 let mut grads = model.new_full_grad_state();
6264 let mut recurrent = model.new_recurrent_grad_state();
6265 let scope = TrainScopeMask {
6266 embed: true,
6267 pre_norm: true,
6268 attn_norm: true,
6269 ffn_norm: true,
6270 attn: true,
6271 ffn: true,
6272 head: true,
6273 bias: false,
6274 };
6275 let grad_scale = 1.0f32 / (steps.len() as f32);
6276 for idx in (0..steps.len()).rev() {
6277 model
6278 .accumulate_token_step_gradients(
6279 &mut scratch,
6280 &traces[idx],
6281 &states[idx + 1],
6282 steps[idx].1,
6283 &pdfs[idx],
6284 grad_scale,
6285 scope,
6286 &mut grads,
6287 None,
6288 &mut recurrent,
6289 )
6290 .expect("segment gradient accumulation");
6291 }
6292 grads
6293 }
6294
6295 #[derive(Clone, Copy, Debug)]
6296 enum Probe {
6297 Embed,
6298 LnOutW,
6299 AttnNormW,
6300 OProj,
6301 KProj,
6302 VProj,
6303 FfnKey,
6304 }
6305
6306 fn probe_value(model: &Model, probe: Probe) -> f32 {
6307 match probe {
6308 Probe::Embed => model.embeddings[7],
6309 Probe::LnOutW => model.ln_out_w[5],
6310 Probe::AttnNormW => model.blocks[0].attn_norm_w[9],
6311 Probe::OProj => model.blocks[0].attn.o_proj[23],
6312 Probe::KProj => model.blocks[0].attn.rkv_proj[64 * 64 + 17],
6313 Probe::VProj => model.blocks[0].attn.rkv_proj[2 * 64 * 64 + 29],
6314 Probe::FfnKey => model.blocks[0].ffn.key_w[11],
6315 }
6316 }
6317
6318 fn set_probe(model: &mut Model, probe: Probe, value: f32) {
6319 match probe {
6320 Probe::Embed => model.embeddings[7] = value,
6321 Probe::LnOutW => model.ln_out_w[5] = value,
6322 Probe::AttnNormW => model.blocks[0].attn_norm_w[9] = value,
6323 Probe::OProj => model.blocks[0].attn.o_proj[23] = value,
6324 Probe::KProj => model.blocks[0].attn.rkv_proj[64 * 64 + 17] = value,
6325 Probe::VProj => model.blocks[0].attn.rkv_proj[2 * 64 * 64 + 29] = value,
6326 Probe::FfnKey => model.blocks[0].ffn.key_w[11] = value,
6327 }
6328 }
6329
6330 fn probe_grad(grads: &FullGradState, probe: Probe) -> f32 {
6331 match probe {
6332 Probe::Embed => grads.embeddings[7],
6333 Probe::LnOutW => grads.ln_out_w[5],
6334 Probe::AttnNormW => grads.blocks[0].attn_norm_w[9],
6335 Probe::OProj => grads.blocks[0].attn.o_proj[23],
6336 Probe::KProj => grads.blocks[0].attn.rkv_proj[64 * 64 + 17],
6337 Probe::VProj => grads.blocks[0].attn.rkv_proj[2 * 64 * 64 + 29],
6338 Probe::FfnKey => grads.blocks[0].ffn.key_w[11],
6339 }
6340 }
6341
6342 fn weighted_checksum(data: &[f32]) -> f64 {
6343 data.iter()
6344 .enumerate()
6345 .map(|(i, &v)| (i as f64 + 1.0) * (v as f64))
6346 .sum()
6347 }
6348
6349 #[test]
6350 fn test_config_default() {
6351 let cfg = Config::default();
6352 assert_eq!(cfg.vocab_size, 256);
6353 assert_eq!(cfg.hidden_size, 256);
6354 assert_eq!(cfg.num_layers, 12);
6355 assert_eq!(cfg.num_heads, 4);
6356 assert_eq!(cfg.head_dim, 64);
6357 }
6358
6359 #[test]
6360 fn validate_rejects_invalid_config_shapes() {
6361 let zero_vocab = Config {
6362 vocab_size: 0,
6363 ..test_cfg()
6364 };
6365 assert!(
6366 zero_vocab
6367 .validate()
6368 .expect_err("zero vocab must fail")
6369 .to_string()
6370 .contains("vocab_size must be > 0")
6371 );
6372
6373 let bad_head_dim = Config {
6374 head_dim: 32,
6375 hidden_size: 32,
6376 ..test_cfg()
6377 };
6378 assert!(
6379 bad_head_dim
6380 .validate()
6381 .expect_err("bad head_dim must fail")
6382 .to_string()
6383 .contains("head_dim must be 64")
6384 );
6385
6386 let bad_hidden = Config {
6387 hidden_size: 128,
6388 ..test_cfg()
6389 };
6390 assert!(
6391 bad_hidden
6392 .validate()
6393 .expect_err("hidden mismatch must fail")
6394 .to_string()
6395 .contains("hidden_size must equal num_heads * head_dim")
6396 );
6397
6398 let zero_layers = Config {
6399 num_layers: 0,
6400 ..test_cfg()
6401 };
6402 assert!(
6403 zero_layers
6404 .validate()
6405 .expect_err("zero layers must fail")
6406 .to_string()
6407 .contains("num_layers must be > 0")
6408 );
6409
6410 let zero_intermediate = Config {
6411 intermediate_size: 0,
6412 ..test_cfg()
6413 };
6414 assert!(
6415 zero_intermediate
6416 .validate()
6417 .expect_err("zero intermediate must fail")
6418 .to_string()
6419 .contains("intermediate_size must be > 0")
6420 );
6421 }
6422
6423 #[test]
6424 fn state_reset_clears_forward_mutation() {
6425 let cfg = test_cfg();
6426 cfg.validate().expect("valid cfg");
6427 let model = Model::new_random(cfg.clone(), 0x5151).expect("random model");
6428 let mut state = model.new_state();
6429 let mut scratch = ScratchBuffers::new(&cfg);
6430
6431 let _ = model.forward(&mut scratch, 7, &mut state);
6432 let _ = model.forward(&mut scratch, 11, &mut state);
6433 assert!(state.v_first_set, "forward pass should initialize v_first");
6434 assert!(
6435 state.layers[0]
6436 .att_state
6437 .as_slice()
6438 .iter()
6439 .any(|&v| v != 0.0),
6440 "forward pass should mutate recurrent state"
6441 );
6442
6443 state.reset();
6444 assert!(!state.v_first_set);
6445 assert!(state.v_first.as_slice().iter().all(|&v| v == 0.0));
6446 for layer in &state.layers {
6447 assert!(layer.att_x_prev.as_slice().iter().all(|&v| v == 0.0));
6448 assert!(layer.att_state.as_slice().iter().all(|&v| v == 0.0));
6449 assert!(layer.ffn_x_prev.as_slice().iter().all(|&v| v == 0.0));
6450 }
6451 }
6452
6453 #[test]
6454 fn train_scope_mask_reports_expected_semantics() {
6455 let none = TrainScopeMask::default();
6456 assert!(!none.trains_non_head_params());
6457 assert!(!none.trains_any_params());
6458
6459 let head_only = TrainScopeMask {
6460 head: true,
6461 ..TrainScopeMask::default()
6462 };
6463 assert!(!head_only.trains_non_head_params());
6464 assert!(head_only.trains_any_params());
6465
6466 let all = TrainScopeMask::all();
6467 assert!(all.embed);
6468 assert!(all.pre_norm);
6469 assert!(all.attn_norm);
6470 assert!(all.ffn_norm);
6471 assert!(all.attn);
6472 assert!(all.ffn);
6473 assert!(all.head);
6474 assert!(all.bias);
6475 assert!(all.trains_non_head_params());
6476 assert!(all.trains_any_params());
6477 }
6478
6479 #[test]
6480 fn save_load_safetensors_roundtrip_preserves_forward_bits() {
6481 let cfg = Config {
6482 num_layers: 2,
6483 intermediate_size: 128,
6484 decay_low_rank: 16,
6485 a_low_rank: 16,
6486 v_low_rank: 16,
6487 g_low_rank: 32,
6488 ..test_cfg()
6489 };
6490 cfg.validate().expect("valid cfg");
6491 let model = Model::new_random(cfg.clone(), 0xBEEF_CAFE).expect("random model");
6492 let path = temp_path("rwkv_roundtrip", "safetensors");
6493 model.save_safetensors(&path).expect("save model");
6494 let loaded = Model::load(&path).expect("load model");
6495
6496 assert_eq!(loaded.config().vocab_size, model.config().vocab_size);
6497 assert_eq!(loaded.config().hidden_size, model.config().hidden_size);
6498 assert_eq!(loaded.config().num_layers, model.config().num_layers);
6499
6500 let mut original_state = model.new_state();
6501 let mut loaded_state = loaded.new_state();
6502 let mut original_scratch = ScratchBuffers::new(&cfg);
6503 let mut loaded_scratch = ScratchBuffers::new(&cfg);
6504 for &token in &[0u32, 7, 31, 99, 255] {
6505 let original_logits = model.forward(&mut original_scratch, token, &mut original_state);
6506 let loaded_logits = loaded.forward(&mut loaded_scratch, token, &mut loaded_state);
6507 for (&lhs, &rhs) in original_logits.iter().zip(loaded_logits.iter()) {
6508 assert_eq!(lhs.to_bits(), rhs.to_bits());
6509 }
6510 }
6511
6512 std::fs::remove_file(path).ok();
6513 }
6514
6515 #[test]
6516 fn save_load_full_adam_roundtrip_preserves_selected_moments() {
6517 let cfg = test_cfg();
6518 cfg.validate().expect("valid cfg");
6519 let model = Model::new_random(cfg, 0xACED).expect("random model");
6520 let mut adam = model.new_full_adam_state();
6521 adam.embeddings.m[0] = 1.25;
6522 adam.embeddings.v[1] = 2.5;
6523 adam.ln_out_w.m[2] = -3.0;
6524 adam.blocks[0].attn.x_r.m[3] = 4.5;
6525 adam.blocks[0].ffn.key_w.v[4] = 5.75;
6526
6527 let path = temp_path("rwkv_adam", "safetensors");
6528 model
6529 .save_full_adam_safetensors(&adam, &path)
6530 .expect("save adam");
6531 let loaded = model.load_full_adam_safetensors(&path).expect("load adam");
6532
6533 assert_eq!(
6534 loaded.embeddings.m[0].to_bits(),
6535 adam.embeddings.m[0].to_bits()
6536 );
6537 assert_eq!(
6538 loaded.embeddings.v[1].to_bits(),
6539 adam.embeddings.v[1].to_bits()
6540 );
6541 assert_eq!(loaded.ln_out_w.m[2].to_bits(), adam.ln_out_w.m[2].to_bits());
6542 assert_eq!(
6543 loaded.blocks[0].attn.x_r.m[3].to_bits(),
6544 adam.blocks[0].attn.x_r.m[3].to_bits()
6545 );
6546 assert_eq!(
6547 loaded.blocks[0].ffn.key_w.v[4].to_bits(),
6548 adam.blocks[0].ffn.key_w.v[4].to_bits()
6549 );
6550
6551 std::fs::remove_file(path).ok();
6552 }
6553
6554 #[test]
6555 fn test_forward_deterministic_snapshot() {
6556 let cfg = Config {
6557 vocab_size: 256,
6558 hidden_size: 64,
6559 num_layers: 2,
6560 num_heads: 1,
6561 head_dim: 64,
6562 intermediate_size: 128,
6563 layer_norm_eps: 1e-5,
6564 group_norm_eps: 64e-5,
6565 decay_low_rank: 16,
6566 a_low_rank: 16,
6567 v_low_rank: 16,
6568 g_low_rank: 32,
6569 };
6570 cfg.validate().expect("valid test config");
6571
6572 let model = Model::new_random(cfg.clone(), 0x1234_5678_9ABC_DEF0).expect("random model");
6573 let mut state = model.new_state();
6574 let mut scratch = ScratchBuffers::new(&cfg);
6575 let tokens = [0u32, 1, 7, 42, 255, 3, 128, 64, 17, 99];
6576
6577 let mut probes = Vec::new();
6578 let mut last_logits = vec![0.0; 8];
6579
6580 for &token in &tokens {
6581 let logits = model.forward(&mut scratch, token, &mut state);
6582 probes.push(logits[0]);
6583 probes.push(logits[1]);
6584 probes.push(logits[2]);
6585 probes.push(logits[42]);
6586 probes.push(logits[127]);
6587 probes.push(logits[255]);
6588 last_logits.copy_from_slice(&logits[0..8]);
6589 }
6590
6591 let probe_checksum = weighted_checksum(&probes);
6592 let last_logits_checksum = weighted_checksum(&last_logits);
6593 let state_att_checksum = weighted_checksum(state.layers[0].att_state.as_slice());
6594 let state_prev_checksum = weighted_checksum(state.layers[1].att_x_prev.as_slice());
6595 let v_first_checksum = weighted_checksum(state.v_first.as_slice());
6596
6597 let expected_probe_checksum = 25.674_967_924_598_604_f64;
6598 let expected_last_logits_checksum = 0.679_873_816_668_987_3_f64;
6599 let expected_state_att_checksum = 129.962_464_237_222_32_f64;
6600 let expected_state_prev_checksum = -231.326_208_570_972_08_f64;
6601 let expected_v_first_checksum = -1.921_361_377_462_744_7_f64;
6602
6603 let tol = 2e-4_f64;
6604 assert!(
6605 (probe_checksum - expected_probe_checksum).abs() <= tol,
6606 "probe_checksum={probe_checksum}"
6607 );
6608 assert!(
6609 (last_logits_checksum - expected_last_logits_checksum).abs() <= tol,
6610 "last_logits_checksum={last_logits_checksum}"
6611 );
6612 assert!(
6613 (state_att_checksum - expected_state_att_checksum).abs() <= tol,
6614 "state_att_checksum={state_att_checksum}"
6615 );
6616 assert!(
6617 (state_prev_checksum - expected_state_prev_checksum).abs() <= tol,
6618 "state_prev_checksum={state_prev_checksum}"
6619 );
6620 assert!(
6621 (v_first_checksum - expected_v_first_checksum).abs() <= tol,
6622 "v_first_checksum={v_first_checksum}"
6623 );
6624 }
6625
6626 #[test]
6627 fn traced_and_untraced_forward_match_exactly() {
6628 let cfg = Config {
6629 vocab_size: 256,
6630 hidden_size: 64,
6631 num_layers: 2,
6632 num_heads: 1,
6633 head_dim: 64,
6634 intermediate_size: 128,
6635 layer_norm_eps: 1e-5,
6636 group_norm_eps: 64e-5,
6637 decay_low_rank: 16,
6638 a_low_rank: 16,
6639 v_low_rank: 16,
6640 g_low_rank: 32,
6641 };
6642 cfg.validate().expect("valid test config");
6643 let model = Model::new_random(cfg.clone(), 0xCAFEBABE).expect("random model");
6644 let mut traced_state = model.new_state();
6645 let mut plain_state = model.new_state();
6646 let mut traced_scratch = ScratchBuffers::new(&cfg);
6647 let mut plain_scratch = ScratchBuffers::new(&cfg);
6648 traced_scratch.set_capture_train_trace(true);
6649 plain_scratch.set_capture_train_trace(false);
6650
6651 let tokens = [3u32, 19, 77, 120, 255, 5, 88, 13, 144, 1, 200];
6652 for &token in &tokens {
6653 let traced_logits = model
6654 .forward(&mut traced_scratch, token, &mut traced_state)
6655 .to_vec();
6656 let plain_logits = model
6657 .forward(&mut plain_scratch, token, &mut plain_state)
6658 .to_vec();
6659 for (a, b) in traced_logits.iter().zip(plain_logits.iter()) {
6660 assert_eq!(a.to_bits(), b.to_bits());
6661 }
6662 assert_eq!(traced_state.v_first_set, plain_state.v_first_set);
6663 for (&a, &b) in traced_state
6664 .v_first
6665 .as_slice()
6666 .iter()
6667 .zip(plain_state.v_first.as_slice())
6668 {
6669 assert_eq!(a.to_bits(), b.to_bits());
6670 }
6671 for (tr_layer, plain_layer) in traced_state.layers.iter().zip(plain_state.layers.iter())
6672 {
6673 for (&a, &b) in tr_layer
6674 .att_x_prev
6675 .as_slice()
6676 .iter()
6677 .zip(plain_layer.att_x_prev.as_slice())
6678 {
6679 assert_eq!(a.to_bits(), b.to_bits());
6680 }
6681 for (&a, &b) in tr_layer
6682 .att_state
6683 .as_slice()
6684 .iter()
6685 .zip(plain_layer.att_state.as_slice())
6686 {
6687 assert_eq!(a.to_bits(), b.to_bits());
6688 }
6689 for (&a, &b) in tr_layer
6690 .ffn_x_prev
6691 .as_slice()
6692 .iter()
6693 .zip(plain_layer.ffn_x_prev.as_slice())
6694 {
6695 assert_eq!(a.to_bits(), b.to_bits());
6696 }
6697 }
6698 }
6699 }
6700
6701 #[test]
6702 fn tbptt_segment_gradients_match_finite_difference() {
6703 let cfg = test_cfg();
6704 cfg.validate().expect("valid test config");
6705 let model = Model::new_random(cfg.clone(), 0xD00D_F00D).expect("random model");
6706 let steps = [(0u32, 1u8), (1, 2), (2, 3)];
6707 let grads = segment_grads(&model, &cfg, &steps);
6708 let eps = 1e-3f32;
6709
6710 for probe in [
6711 Probe::Embed,
6712 Probe::LnOutW,
6713 Probe::AttnNormW,
6714 Probe::OProj,
6715 Probe::KProj,
6716 Probe::VProj,
6717 Probe::FfnKey,
6718 ] {
6719 let analytic = probe_grad(&grads, probe);
6720
6721 let mut plus = model.clone();
6722 let base = probe_value(&plus, probe);
6723 set_probe(&mut plus, probe, base + eps);
6724 let loss_plus = segment_loss(&plus, &cfg, &steps);
6725
6726 let mut minus = model.clone();
6727 set_probe(&mut minus, probe, base - eps);
6728 let loss_minus = segment_loss(&minus, &cfg, &steps);
6729
6730 let numeric = -((loss_plus - loss_minus) / (2.0 * eps as f64)) as f32;
6731 let tol = 5e-2f32.max(analytic.abs().max(numeric.abs()) * 8e-2);
6732 assert!(
6733 (analytic - numeric).abs() <= tol,
6734 "probe={probe:?} analytic={analytic} numeric={numeric} tol={tol}"
6735 );
6736 }
6737 }
6738
6739 #[test]
6740 fn tbptt_sgd_step_reduces_mean_segment_loss() {
6741 let cfg = test_cfg();
6742 cfg.validate().expect("valid test config");
6743 let mut model = Model::new_random(cfg.clone(), 0x1234_5678).expect("random model");
6744 let steps = [(0u32, 1u8), (1, 2), (2, 3), (3, 4)];
6745 let before = segment_loss(&model, &cfg, &steps);
6746
6747 let mut scratch = ScratchBuffers::new(&cfg);
6748 let mut workspace = TbpttReplayWorkspace::new(&model);
6749 let start_state = model.new_state();
6750 let mut live_state = model.new_state();
6751 let mut adam_t = 0usize;
6752 let scope = TrainScopeMask {
6753 embed: true,
6754 pre_norm: true,
6755 attn_norm: true,
6756 ffn_norm: true,
6757 attn: true,
6758 ffn: true,
6759 head: true,
6760 bias: false,
6761 };
6762
6763 model
6764 .online_train_segment_tbptt(
6765 &mut scratch,
6766 &mut workspace,
6767 &start_state,
6768 &steps,
6769 scope,
6770 OptimizerKind::Sgd,
6771 1e-3,
6772 0.0,
6773 2,
6774 &mut adam_t,
6775 None,
6776 None,
6777 None,
6778 None,
6779 &mut live_state,
6780 )
6781 .expect("tbptt sgd step");
6782
6783 let after = segment_loss(&model, &cfg, &steps);
6784 assert!(
6785 after < before,
6786 "expected SGD TBPTT step to reduce mean loss: before={before} after={after}"
6787 );
6788 }
6789
6790 #[test]
6791 fn head_only_bptt1_update_succeeds_without_full_trace() {
6792 let cfg = test_cfg();
6793 cfg.validate().expect("valid cfg");
6794 let mut model = Model::new_random(cfg.clone(), 0xAAAA_5555).expect("random model");
6795 let mut scratch = ScratchBuffers::new(&cfg);
6796 let mut state = model.new_state();
6797 scratch.set_capture_train_trace(false);
6798
6799 let logits = model.forward(&mut scratch, 9, &mut state).to_vec();
6800 let mut pdf = vec![0.0f64; cfg.vocab_size];
6801 super::super::super::softmax_pdf_floor_with_bias(&logits, None, &mut pdf);
6802 let before = model.lm_head_weights()[0];
6803 let mut adam_t = 0usize;
6804 let scope = TrainScopeMask {
6805 head: true,
6806 ..TrainScopeMask::default()
6807 };
6808
6809 model
6810 .online_train_step_bptt1(
6811 &mut scratch,
6812 &state,
6813 7,
6814 &pdf,
6815 scope,
6816 OptimizerKind::Sgd,
6817 1e-3,
6818 0.0,
6819 &mut adam_t,
6820 None,
6821 None,
6822 None,
6823 None,
6824 )
6825 .expect("head-only update");
6826
6827 assert_ne!(model.lm_head_weights()[0].to_bits(), before.to_bits());
6828 }
6829
6830 #[test]
6831 fn full_training_bptt1_requires_captured_trace() {
6832 let cfg = test_cfg();
6833 cfg.validate().expect("valid cfg");
6834 let mut model = Model::new_random(cfg.clone(), 0x1234).expect("random model");
6835 let mut scratch = ScratchBuffers::new(&cfg);
6836 let mut state = model.new_state();
6837 scratch.set_capture_train_trace(false);
6838
6839 let logits = model.forward(&mut scratch, 4, &mut state).to_vec();
6840 let mut pdf = vec![0.0f64; cfg.vocab_size];
6841 super::super::super::softmax_pdf_floor_with_bias(&logits, None, &mut pdf);
6842
6843 let err = model
6844 .online_train_step_bptt1(
6845 &mut scratch,
6846 &state,
6847 3,
6848 &pdf,
6849 TrainScopeMask {
6850 attn: true,
6851 ..TrainScopeMask::default()
6852 },
6853 OptimizerKind::Sgd,
6854 1e-3,
6855 0.0,
6856 &mut 0usize,
6857 None,
6858 None,
6859 None,
6860 None,
6861 )
6862 .expect_err("non-head training should require trace");
6863 assert!(err.to_string().contains("full training trace is missing"));
6864 }
6865
6866 #[test]
6867 fn adam_full_training_requires_explicit_adam_state() {
6868 let cfg = test_cfg();
6869 cfg.validate().expect("valid cfg");
6870 let mut model = Model::new_random(cfg.clone(), 0xCAFE).expect("random model");
6871 let mut scratch = ScratchBuffers::new(&cfg);
6872 let mut state = model.new_state();
6873 scratch.set_capture_train_trace(true);
6874
6875 let logits = model.forward(&mut scratch, 5, &mut state).to_vec();
6876 let mut pdf = vec![0.0f64; cfg.vocab_size];
6877 super::super::super::softmax_pdf_floor_with_bias(&logits, None, &mut pdf);
6878
6879 let err = model
6880 .online_train_step_bptt1(
6881 &mut scratch,
6882 &state,
6883 6,
6884 &pdf,
6885 TrainScopeMask {
6886 attn: true,
6887 ..TrainScopeMask::default()
6888 },
6889 OptimizerKind::Adam,
6890 1e-3,
6891 0.0,
6892 &mut 0usize,
6893 None,
6894 None,
6895 None,
6896 None,
6897 )
6898 .expect_err("adam full training should require optimizer state");
6899 assert!(
6900 err.to_string()
6901 .contains("Adam full-training state is missing")
6902 );
6903 }
6904}