refactor(ml-alpha): move LN_a weights from PerceptionTrainer to CfcTrunk (X4)

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).
This commit is contained in:
jgrusewski
2026-05-19 01:41:19 +02:00
parent 2d849dd5e3
commit 868021e818

View File

@@ -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<f32>,
ln_a_bias_d: CudaSlice<f32>,
// X4: LN_a weights moved to `self.trunk.ln_a_gain_d` / `ln_a_bias_d`.
ln_a_stats_d: CudaSlice<f32>,
/// `[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);