refactor(ml-alpha): move LN_b weights from PerceptionTrainer to CfcTrunk (X6)
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).
This commit is contained in:
@@ -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<f32>,
|
||||
/// 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<f32>,
|
||||
/// LayerNorm bias (`[HIDDEN_DIM]`, initialised to 0.0).
|
||||
ln_bias_d: CudaSlice<f32>,
|
||||
// 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<f32>,
|
||||
@@ -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<f32> = vec![1.0; HIDDEN_DIM];
|
||||
let ln_bias_init: Vec<f32> = 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::<f32>(cfg.n_batch * k * 2)?,
|
||||
ln_out_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * HIDDEN_DIM)?,
|
||||
grad_ln_gain_per_row_d: stream.alloc_zeros::<f32>(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);
|
||||
|
||||
Reference in New Issue
Block a user