From 335769943102238be468bc697ae62ec243cda538 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 19 May 2026 01:44:48 +0200 Subject: [PATCH] refactor(ml-alpha): move LN_b weights from PerceptionTrainer to CfcTrunk (X6) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit LN_b (formerly ln_gain_d / ln_bias_d on trainer) now lives at self.trunk.ln_b_gain_d / ln_b_bias_d. LN_b is the LayerNorm after Mamba2 stack 2. Verification: golden bit-exact; ml-alpha lib green. Per spec §2.2 (X6). --- crates/ml-alpha/src/trainer/perception.rs | 34 +++++++++++------------ 1 file changed, 17 insertions(+), 17 deletions(-) diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index 6d18f1d73..e68e920e8 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -297,12 +297,9 @@ pub struct PerceptionTrainer { /// tracks to compute dynamic per-horizon weights. Refreshed by /// the BCE kernel on every training step. Shape: [N_HORIZONS]. loss_per_horizon_d: CudaSlice, - /// LayerNorm gain (`[HIDDEN_DIM]`, initialised to 1.0). Applied - /// to Mamba2's `h_enriched_seq` output before the K-loop CfC - /// consumes it. Phase 1: stabilises the trunk-to-CfC distribution. - ln_gain_d: CudaSlice, - /// LayerNorm bias (`[HIDDEN_DIM]`, initialised to 0.0). - ln_bias_d: CudaSlice, + // X6: LN_b weights (formerly ln_gain_d / ln_bias_d) moved to + // `self.trunk.ln_b_gain_d` / `ln_b_bias_d`. LN_b is the LayerNorm + // AFTER Mamba2 stack 2 — stabilises the trunk-to-CfC distribution. /// Per-row LN stats `[B*K, 2]` (mean, inv_std) saved by fwd, used /// by bwd. ln_stats_d: CudaSlice, @@ -799,8 +796,13 @@ impl PerceptionTrainer { // Two instances: LN_a between m1 and m2, LN_b between m2 and CfC. let ln_gain_init: Vec = vec![1.0; HIDDEN_DIM]; let ln_bias_init: Vec = vec![0.0; HIDDEN_DIM]; - let ln_gain_d = upload(&stream, &ln_gain_init)?; - let ln_bias_d = upload(&stream, &ln_bias_init)?; + // X6: upload LN_b weights into the trunk's slots. + stream + .memcpy_htod(&ln_gain_init, &mut trunk.ln_b_gain_d) + .context("trunk.ln_b_gain_d upload")?; + stream + .memcpy_htod(&ln_bias_init, &mut trunk.ln_b_bias_d) + .context("trunk.ln_b_bias_d upload")?; let mut opt_ln_gain = AdamW::new(dev, HIDDEN_DIM, cfg.lr_cfc)?; opt_ln_gain.wd = 0.0; let mut opt_ln_bias = AdamW::new(dev, HIDDEN_DIM, cfg.lr_cfc)?; @@ -911,8 +913,6 @@ impl PerceptionTrainer { // HIDDEN_DIM (128) by construction — checked at kernel // launch time. Per-row scratch is [B*K, HIDDEN]; the // single param-grad reducer collapses it to [HIDDEN]. - ln_gain_d, - ln_bias_d, ln_stats_d: stream.alloc_zeros::(cfg.n_batch * k * 2)?, ln_out_d: stream.alloc_zeros::(cfg.n_batch * k * HIDDEN_DIM)?, grad_ln_gain_per_row_d: stream.alloc_zeros::(cfg.n_batch * k * HIDDEN_DIM)?, @@ -1483,8 +1483,8 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_fwd_fn); launch .arg(self.mamba2_l2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_gain_d) - .arg(&self.ln_bias_d) + .arg(&self.trunk.ln_b_gain_d) + .arg(&self.trunk.ln_b_bias_d) .arg(&n_rows_ln) .arg(&mut self.ln_out_d) .arg(&mut self.ln_stats_d); @@ -2002,7 +2002,7 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_bwd_fn); launch .arg(self.mamba2_l2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_gain_d) + .arg(&self.trunk.ln_b_gain_d) .arg(&self.ln_stats_d) .arg(self.grad_h_enriched_seq_d.cuda_data()) .arg(&n_rows_ln) @@ -2285,8 +2285,8 @@ impl PerceptionTrainer { self.opt_heads_b_main.step(&mut self.heads_b_main_d, &self.grad_heads_b_main_d)?; self.opt_heads_w_skip.step(&mut self.heads_w_skip_d, &self.grad_heads_w_skip_d)?; 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_gain.step(&mut self.trunk.ln_b_gain_d, &self.grad_ln_gain_d)?; + self.opt_ln_bias.step(&mut self.trunk.ln_b_bias_d, &self.grad_ln_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)?; @@ -2523,8 +2523,8 @@ impl PerceptionTrainer { let mut launch = self.stream.launch_builder(&self.ln_fwd_fn); launch .arg(self.mamba2_l2_fwd_scratch.h_enriched_seq.cuda_data()) - .arg(&self.ln_gain_d) - .arg(&self.ln_bias_d) + .arg(&self.trunk.ln_b_gain_d) + .arg(&self.trunk.ln_b_bias_d) .arg(&n_rows_ln) .arg(&mut self.ln_out_d) .arg(&mut self.ln_stats_d);