diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index 57b9ee683..af28441eb 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -168,8 +168,7 @@ pub struct PerceptionTrainer { // ── LN_a (between Mamba2 stacks; Phase 2B) ── // The existing LN sits between m2 and CfC (now called LN_b). LN_a // is a separate instance with its own learnable gain/bias. - ln_a_gain_d: CudaSlice, - ln_a_bias_d: CudaSlice, + // X4: LN_a weights moved to `self.trunk.ln_a_gain_d` / `ln_a_bias_d`. ln_a_stats_d: CudaSlice, /// `[B, K, HIDDEN_DIM]` — LN_a forward output, fed as input to m2. ln_a_out_d: GpuTensor, @@ -806,8 +805,13 @@ impl PerceptionTrainer { opt_ln_gain.wd = 0.0; let mut opt_ln_bias = AdamW::new(dev, HIDDEN_DIM, cfg.lr_cfc)?; opt_ln_bias.wd = 0.0; - let ln_a_gain_d = upload(&stream, &ln_gain_init)?; - let ln_a_bias_d = upload(&stream, &ln_bias_init)?; + // X4: upload LN_a weights into the trunk's slots. + stream + .memcpy_htod(&ln_gain_init, &mut trunk.ln_a_gain_d) + .context("trunk.ln_a_gain_d upload")?; + stream + .memcpy_htod(&ln_bias_init, &mut trunk.ln_a_bias_d) + .context("trunk.ln_a_bias_d upload")?; let mut opt_ln_a_gain = AdamW::new(dev, HIDDEN_DIM, cfg.lr_cfc)?; opt_ln_a_gain.wd = 0.0; let mut opt_ln_a_bias = AdamW::new(dev, HIDDEN_DIM, cfg.lr_cfc)?; @@ -1073,8 +1077,6 @@ impl PerceptionTrainer { mamba2_l2_bwd_scratch, mamba2_l2_grads_buffers, // LN_a between Mamba2 stacks. - ln_a_gain_d, - ln_a_bias_d, ln_a_stats_d, ln_a_out_d, grad_ln_a_gain_per_row_d, @@ -1454,8 +1456,8 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_fwd_fn); launch .arg(self.mamba2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_a_gain_d) - .arg(&self.ln_a_bias_d) + .arg(&self.trunk.ln_a_gain_d) + .arg(&self.trunk.ln_a_bias_d) .arg(&n_rows_ln) .arg(self.ln_a_out_d.data_mut()) .arg(&mut self.ln_a_stats_d); @@ -2075,7 +2077,7 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_bwd_fn); launch .arg(self.mamba2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_a_gain_d) + .arg(&self.trunk.ln_a_gain_d) .arg(&self.ln_a_stats_d) .arg(self.mamba2_l2_grads_buffers.d_x_from_in.cuda_data()) .arg(&n_rows_ln) @@ -2284,8 +2286,8 @@ impl PerceptionTrainer { self.opt_heads_b_skip.step(&mut self.heads_b_skip_d, &self.grad_heads_b_skip_d)?; self.opt_ln_gain.step(&mut self.ln_gain_d, &self.grad_ln_gain_d)?; self.opt_ln_bias.step(&mut self.ln_bias_d, &self.grad_ln_bias_d)?; - self.opt_ln_a_gain.step(&mut self.ln_a_gain_d, &self.grad_ln_a_gain_d)?; - self.opt_ln_a_bias.step(&mut self.ln_a_bias_d, &self.grad_ln_a_bias_d)?; + self.opt_ln_a_gain.step(&mut self.trunk.ln_a_gain_d, &self.grad_ln_a_gain_d)?; + self.opt_ln_a_bias.step(&mut self.trunk.ln_a_bias_d, &self.grad_ln_a_bias_d)?; self.opt_vsn_w.step(&mut self.trunk.vsn_w_d, &self.grad_vsn_w_d)?; self.opt_vsn_b.step(&mut self.trunk.vsn_b_d, &self.grad_vsn_b_d)?; self.opt_attn_q.step(&mut self.attn_q_d, &self.grad_attn_q_d)?; @@ -2493,8 +2495,8 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_fwd_fn); launch .arg(self.mamba2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_a_gain_d) - .arg(&self.ln_a_bias_d) + .arg(&self.trunk.ln_a_gain_d) + .arg(&self.trunk.ln_a_bias_d) .arg(&n_rows_ln) .arg(self.ln_a_out_d.data_mut()) .arg(&mut self.ln_a_stats_d);