Skip to main content

infotheory/backends/rwkvzip/rwkv7/
model.rs

1//! RWKV7 model implementation with portable SIMD-optimized inference.
2//!
3//! This is a high-performance implementation built on `wide` kernels.
4//! Single-token inference is the primary use case (streaming compression).
5
6use 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/// Model configuration.
21#[derive(Debug, Clone)]
22pub struct Config {
23    /// Vocabulary size (byte-level models use 256).
24    pub vocab_size: usize,
25    /// Hidden channel width `C`.
26    pub hidden_size: usize,
27    /// Number of RWKV blocks.
28    pub num_layers: usize,
29    /// Number of attention heads `H`.
30    pub num_heads: usize,
31    /// Per-head channel width `N` (currently fixed to 64 in kernels).
32    pub head_dim: usize,
33    /// Feed-forward intermediate width.
34    pub intermediate_size: usize,
35    /// Epsilon for layer normalization.
36    pub layer_norm_eps: f32,
37    /// Epsilon for group normalization (`64e-5` in reference).
38    pub group_norm_eps: f32,
39
40    /// Low-rank width for decay projection.
41    pub decay_low_rank: usize, // w_lora
42    /// Low-rank width for `a` projection.
43    pub a_low_rank: usize,
44    /// Low-rank width for `v` projection.
45    pub v_low_rank: usize,
46    /// Low-rank width for `g` projection.
47    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, // 256 / 64
57            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    /// Validate configuration invariants required by current kernels.
71    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/// Per-layer state for RWKV7.
97#[derive(Clone)]
98pub struct LayerState {
99    /// Previous token embedding for attention time-shift (hidden_size,)
100    pub att_x_prev: Tensor1D,
101    /// Attention state matrix (num_heads, head_dim, head_dim) = (H, N, N)
102    pub att_state: Tensor1D, // Flat for SIMD access
103    /// Previous token embedding for FFN time-shift (hidden_size,)
104    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/// Full model state.
125#[derive(Clone)]
126pub struct State {
127    /// Per-layer recurrent state.
128    pub layers: Vec<LayerState>,
129    /// First layer's value output (for residual connection) - pre-allocated
130    pub v_first: Tensor1D,
131    /// Flag to indicate if v_first has been set
132    pub v_first_set: bool,
133}
134
135impl State {
136    /// Allocate a zero-initialized recurrent state for a model configuration.
137    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    /// Reset recurrent buffers to their initial (all-zero) state.
146    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/// Weights for a single attention layer.
167#[derive(Clone)]
168struct AttentionWeights {
169    // Token shift mixing factors
170    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    // Packed r/k/v projections for parallel computation
178    // Layout: [r_proj (C*C), k_proj (C*C), v_proj (C*C)]
179    rkv_proj: Tensor1D,
180
181    // Output projection (stored transposed for efficient gemv)
182    o_proj: Tensor1D,
183
184    // Low-rank W: w = tanh(x @ w1) @ w2 + w0
185    w1: Tensor1D, // (C, D_w)
186    w2: Tensor1D, // (D_w, C)
187    w0: Tensor1D, // (C,)
188
189    // Low-rank A: a = sigmoid(x @ a1 @ a2 + a0)
190    a1: Tensor1D, // (C, D_a)
191    a2: Tensor1D, // (D_a, C)
192    a0: Tensor1D, // (C,)
193
194    // Low-rank V (layers > 0): nu = sigmoid(x @ v1 @ v2 + v0)
195    v1: Option<Tensor1D>, // (C, D_v)
196    v2: Option<Tensor1D>, // (D_v, C)
197    v0: Option<Tensor1D>, // (C,)
198
199    // Low-rank G: g = sigmoid(x @ g1) @ g2
200    g1: Tensor1D, // (C, D_g)
201    g2: Tensor1D, // (D_g, C)
202
203    // Key scaling
204    k_k: Tensor1D, // (C,)
205    k_a: Tensor1D, // (C,)
206    r_k: Tensor1D, // (H, N)
207
208    // Group norm for output
209    g_norm_w: Tensor1D, // (C,)
210    g_norm_b: Tensor1D, // (C,)
211}
212
213/// Weights for a single FFN layer.
214#[derive(Clone)]
215struct FfnWeights {
216    x_k: Tensor1D,     // (C,) time shift mix
217    key_w: Tensor1D,   // (C, I) -> relu(x @ W)^2
218    value_w: Tensor1D, // (I, C)
219}
220
221/// Weights for a single block.
222#[derive(Clone)]
223struct BlockWeights {
224    // Pre-norm (layer 0 only)
225    pre_norm_w: Option<Tensor1D>,
226    pre_norm_b: Option<Tensor1D>,
227
228    // Attention norm
229    attn_norm_w: Tensor1D,
230    attn_norm_b: Tensor1D,
231
232    // FFN norm
233    ffn_norm_w: Tensor1D,
234    ffn_norm_b: Tensor1D,
235
236    attn: AttentionWeights,
237    ffn: FfnWeights,
238}
239
240/// RWKV7 model.
241#[derive(Clone)]
242pub struct Model {
243    cfg: Config,
244
245    // Embeddings (vocab_size, hidden_size)
246    embeddings: Tensor1D,
247
248    // Output norm
249    ln_out_w: Tensor1D,
250    ln_out_b: Tensor1D,
251
252    // LM head (vocab_size, hidden_size)
253    lm_head: Tensor1D,
254
255    // Layers
256    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)]
323/// Adam moments for full-parameter RWKV online training.
324pub 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)]
333/// Train-scope mask for RWKV full-parameter online updates.
334pub struct TrainScopeMask {
335    /// Train token embeddings.
336    pub embed: bool,
337    /// Train optional pre-norm parameters.
338    pub pre_norm: bool,
339    /// Train attention norm parameters.
340    pub attn_norm: bool,
341    /// Train FFN norm parameters.
342    pub ffn_norm: bool,
343    /// Train attention block parameters.
344    pub attn: bool,
345    /// Train FFN block parameters.
346    pub ffn: bool,
347    /// Train LM-head weights.
348    pub head: bool,
349    /// Train additive output-bias terms.
350    pub bias: bool,
351}
352
353impl TrainScopeMask {
354    #[inline]
355    /// Enable all train scopes.
356    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    /// Returns whether any non-head model parameters are trainable.
371    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    /// Returns whether any parameter/bias updates are enabled.
377    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/// Pre-allocated scratch buffers to avoid allocations in hot path.
675#[derive(Clone)]
676pub struct ScratchBuffers {
677    x: Tensor1D,          // Current hidden state
678    x_normed: Tensor1D,   // After layer norm
679    xr: Tensor1D,         // Token-shifted for r
680    xw: Tensor1D,         // Token-shifted for w
681    xk: Tensor1D,         // Token-shifted for k
682    xv: Tensor1D,         // Token-shifted for v
683    xa: Tensor1D,         // Token-shifted for a
684    xg: Tensor1D,         // Token-shifted for g
685    r: Tensor1D,          // Receptance
686    k: Tensor1D,          // Key
687    v: Tensor1D,          // Value
688    w_lora_tmp: Tensor1D, // Low-rank temp
689    w_decay: Tensor1D,    // Decay factor
690    a: Tensor1D,          // Gate a
691    g: Tensor1D,          // Gate g
692    kk: Tensor1D,         // Normalized key
693    y: Tensor1D,          // WKV output
694    att_out: Tensor1D,    // Attention output
695    ffn_k: Tensor1D,      // FFN key
696    ffn_out: Tensor1D,    // FFN output
697    logits: Tensor1D,     // Output logits
698    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    /// Allocate reusable per-token scratch buffers sized for `cfg`.
757    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    /// Final normalized hidden state consumed by LM head.
820    #[inline]
821    pub fn lm_head_input(&self) -> &[f32] {
822        self.x_normed.as_slice()
823    }
824
825    #[inline]
826    /// Borrow logits from the latest forward pass scratch buffer.
827    pub fn logits(&self) -> &[f32] {
828        self.logits.as_slice()
829    }
830
831    /// Restore LM-head input snapshot for reversible online updates.
832    #[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    /// Enable or disable per-token training trace capture.
838    #[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    /// Whether the current scratch contains a valid full-trace for the latest forward pass.
847    #[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    /// Load model from safetensors file.
865    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        // Infer config from weights
874        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; // Assume head_dim=64
879        let head_dim = 64;
880
881        // Count layers by looking for layer weights
882        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        // Get intermediate size from FFN
891        let ffn_key = weights.require("model.layers.0.ffn.key.weight")?;
892        let intermediate_size = ffn_key.shape()[0];
893
894        // Get low-rank dimensions
895        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        // v_low_rank from layer 1 (layer 0 doesn't have it)
905        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        // Load embeddings
931        let embeddings = Self::tensor_from(&weights, "model.embeddings.weight")?;
932
933        // Load output norm
934        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        // Load LM head
938        let lm_head = Self::tensor_from(&weights, "lm_head.weight")?;
939
940        // Load blocks
941        let mut blocks = Vec::with_capacity(num_layers);
942        for i in 0..num_layers {
943            let prefix = format!("model.layers.{}", i);
944
945            // Pre-norm (layer 0 only)
946            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            // Norms
962            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            // Attention weights
968            // Load r/k/v projections and pack them contiguously
969            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            // Create packed RKV tensor: [r_proj, k_proj, v_proj]
980            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            // FFN weights
1030            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    /// Create a randomly initialized model for online-training workflows.
1059    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    /// Save model weights to a `.safetensors` file plus JSON sidecar config.
1229    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    /// Allocate zero-initialized Adam moments matching all trainable tensors.
1478    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    /// Allocate zero-initialized gradient storage matching all trainable tensors.
1531    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    /// Save full-parameter Adam moments for exact online-training continuation.
1588    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    /// Load full-parameter Adam moments and validate tensor shapes.
1741    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    /// Get model configuration.
1839    pub fn config(&self) -> &Config {
1840        &self.cfg
1841    }
1842
1843    /// Create new state for this model.
1844    pub fn new_state(&self) -> State {
1845        State::new(&self.cfg)
1846    }
1847
1848    /// Immutable LM-head weights, row-major `(vocab, hidden)`.
1849    #[inline]
1850    pub fn lm_head_weights(&self) -> &[f32] {
1851        self.lm_head.as_slice()
1852    }
1853
1854    /// Mutable LM-head weights, row-major `(vocab, hidden)`.
1855    #[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    /// Run one TBPTT training segment and write the resulting live state.
3258    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    /// Perform one exact bptt=1 online training step over the latest forward trace.
3398    #[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            // FFN residual split: x_out = x_after_attn + ffn_out.
3594            scratch.grad_x2.copy_from_slice(scratch.grad_x.as_slice()); // d x_after_attn
3595            scratch.grad_x3.copy_from_slice(scratch.grad_x.as_slice()); // d ffn_out
3596
3597            // ffn_out = value_w @ ffn_k
3598            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            // relu^2 backward
3636            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            // key_w backward
3646            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            // token_shift backward (ffn)
3684            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); // d ffn_norm
3690                scratch.grad_param[col] = g * (prev - base); // d x_k
3691            }
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            // ffn norm backward: d x_after_attn contribution.
3715            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            // Attention residual split.
3764            scratch.grad_x.copy_from_slice(scratch.grad_x2.as_slice()); // d x_after_pre
3765            scratch.grad_x3.copy_from_slice(scratch.grad_x2.as_slice()); // d att_out
3766
3767            // out_proj backward
3768            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            // Gate backward.
3806            for col in 0..c {
3807                let gy = scratch.grad_x4[col];
3808                scratch.grad_saved[col] = gy * tr.y_head[col]; // d g
3809                scratch.grad_x4[col] = gy * tr.g[col]; // d y_head
3810            }
3811
3812            // Head-qk branch.
3813            scratch.grad_x2.zero(); // d r
3814            scratch.grad_x3.zero(); // d k_scaled
3815            scratch.grad_x6.zero(); // d v_final
3816            scratch.grad_param.zero(); // d r_k
3817            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            // GroupNorm backward.
3859            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(),     // d y_wkv
3868                scratch.grad_param.as_mut_slice(),  // d g_norm_w
3869                scratch.grad_param2.as_mut_slice(), // d g_norm_b
3870            );
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            // WKV kernel backward.
3908            scratch.grad_param.zero(); // d w_decay
3909            scratch.grad_x5.zero(); // d a
3910            scratch.grad_param2.zero(); // d kk
3911            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            // k scaling + kk normalization backward.
3993            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; // d k_pre (scaled path)
3998                scratch.grad_x5[col] += d_scale * block.attn.k_a[col]; // d a
3999                scratch.grad_param[col] = d_scale * (tr.a[col] - 1.0); // d k_a
4000            }
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]; // d k_pre from kk_pre
4014                scratch.grad_param2[col] = g * tr.k_pre[col]; // d k_k
4015            }
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            // V-residual backward (layers > 0).
4053            scratch
4054                .grad_param2
4055                .copy_from_slice(scratch.grad_x6.as_slice()); // d v_final snapshot
4056            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); // d v_pre
4068                    scratch.grad_x3[col] = gv * (scratch.train_v_first[col] - tr.v_pre[col]); // d nu
4069                    scratch.grad_v_first[col] += gv * nu; // d v_first
4070                }
4071                for col in 0..c {
4072                    let nu = tr.nu[col];
4073                    scratch.grad_x3[col] *= nu * (1.0 - nu); // d nu_pre
4074                }
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; // add into d xv after projection transpose
4162                }
4163            }
4164
4165            // R/K/V projection updates and input grads.
4166            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            // W low-rank backward.
4261            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); // d w_pre
4266            }
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            // A low-rank backward.
4364            for col in 0..c {
4365                let a = tr.a[col];
4366                scratch.grad_x5[col] *= a * (1.0 - a); // d a_pre
4367            }
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            // G low-rank backward.
4461            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            // Token-shift backward for attention branches.
4539            scratch.grad_x3.zero(); // d attn_norm
4540
4541            // x_r
4542            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            // x_w
4573            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            // x_k
4604            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            // x_v
4635            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            // x_a
4666            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            // x_g
4697            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            // attn norm backward.
4728            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            // pre_norm backward (layer 0 only).
4777            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    /// Forward pass for a single token.
4841    /// Returns logits for next token prediction.
4842    #[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    /// Forward pass that records per-layer timings through a custom sink.
4854    #[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        // Get token embedding
4894        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            // Process each layer (using index to avoid borrow conflicts)
4908            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                // Pre-norm (layer 0 only)
4915                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                // Attention norm
4935                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                // Add attention residual: x = x + att_out
4963                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                // FFN norm
4976                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                // Add FFN residual: x = x + ffn_out
5009                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            // Output norm
5023            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            // LM head: logits = x @ lm_head.T
5033            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        // Update prev state for next token
5098        kernel::copy(
5099            scratch.x_normed.as_ptr(),
5100            layer_state.att_x_prev.as_mut_ptr(),
5101            c,
5102        );
5103
5104        // r/k/v projections from packed matrix (sequential for better cache)
5105        // Packed layout: [r_proj (C*C), k_proj (C*C), v_proj (C*C)]
5106        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        // w decay: w = exp(-sigmoid(tanh(xw @ w1) @ w2 + w0) / sqrt(e))
5136        // Step 1: tmp = xw @ w1.T (D_w output)
5137        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        // Step 2: tanh
5145        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        // Step 3: tmp @ w2.T + w0
5156        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        // Add bias w0
5164        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        // Step 4: exp(-sigmoid(x) / sqrt(e))
5175        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        // a = sigmoid(xa @ a1.T @ a2.T + a0)
5193        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        // g = sigmoid(xg @ g1.T) @ g2.T
5225        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        // Value residual (layer > 0)
5255        if layer_idx == 0 {
5256            // Copy v to v_first buffer (no allocation)
5257            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            // nu = sigmoid(xv @ v1.T @ v2.T + v0)
5269            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(), // reuse as temp
5285                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            // v = v + (v_first - v) * nu
5301            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        // kk = k * k_k, then L2 normalize per head
5316        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        // Normalize per head
5327        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        // k = k * (1 + (a - 1) * k_a)
5342        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        // WKV state update: S = S*w.T - S@kk*(kk*a).T + v*k.T; y = S@r
5352        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        // Group norm
5370        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        // Add head-qk term: y += ((r * k * r_k).sum_per_head) * v
5385        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        // Apply gate: y = y * g
5405        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        // Output projection: att_out = o_proj @ y
5417        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        // Token shift: xk = x_normed + x_k * (prev - x_normed)
5447        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        // Update prev state
5460        kernel::copy(
5461            scratch.x_normed.as_ptr(),
5462            layer_state.ffn_x_prev.as_mut_ptr(),
5463            c,
5464        );
5465
5466        // k = relu(xk @ key_w.T)^2
5467        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        // ffn_out = k @ value_w.T
5485        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}