From 868021e818f51cab3705df1b452c2215c29ee4a3 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 19 May 2026 01:41:19 +0200 Subject: [PATCH] refactor(ml-alpha): move LN_a weights from PerceptionTrainer to CfcTrunk (X4) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit LN_a gain/bias tensors now live at self.trunk.ln_a_gain_d / ln_a_bias_d. Trainer-side initialization values still drive the upload, preserving PRNG-driven init values. Gradient buffers + AdamW state remain at trainer level. Verification: golden bit-exact (max_diff = 0.000000); ml-alpha lib green. Per spec 2026-05-19-ml-alpha-v2-trunk-grows-and-deployability-design.md §2.2 (X4). --- crates/ml-alpha/src/trainer/perception.rs | 28 ++++++++++++----------- 1 file changed, 15 insertions(+), 13 deletions(-) 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);