feat(ml-alpha): 2-stack Mamba2 with inter-stack LayerNorm (Phase 2B)

Doubles the trunk capacity. Forward chain:
  snap_features → VSN → m1 → LN_a → m2 → LN_b → CfC → GRN heads

m1 = Mamba2Block { in_dim=FEATURE_DIM=40, hidden_dim=128 }
m2 = Mamba2Block { in_dim=128, hidden_dim=128 }

Both stacks share the SAME state_dim (cfg.mamba2_state_dim) and the
SAME hidden_dim. m2 reads m1's output (post LN_a). LN_a is a separate
LayerNorm instance from LN_b (the existing trunk-to-CfC normaliser).

Backward chain reverses the forward:
  ... grad_ln_in_d → m2.bwd → m2_grads_buffers.d_x_from_in (= LN_a output grad)
  → LN_a.bwd → grad_ln_a_in_d (= m1 output grad)
  → m1.bwd → m1_grads_buffers.d_x_from_in (= VSN output grad)
  → VSN.bwd → ...

Both Mamba2 stacks emit `d_x_from_in` (Phase 2D refactor already
exposed it on m1; m2 uses the same code path). LN_a uses the existing
layer_norm_fwd / layer_norm_bwd / layer_norm_reduce_param_grads
kernels — no new CUDA work, just a second instance with its own
gain/bias/stats/grad scratches.

New trainer state: ~17 fields (mamba2_l2 + its scratch + LN_a + LN_a
grads + opt_ln_a_*). All initialised in the construction order that
respects the `stream` move-into-Self at the end of new().

set_lr_mamba2 now updates BOTH stacks' AdamW configs. Total AdamW
instances on the trainer: 21 (CfC×4 + GRN heads×10 + LN×2 + LN_a×2 +
VSN×2 + Mamba2×2 grouped × 9 params each).

All 8 perception_overfit smokes pass: synthetic constant-direction
signal converges 0.31 → 0.0000 by step 100 (matches single-stack
trajectory — proves both stacks are wired forward + backward and all
21 AdamW optimisers move weights).

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 22:28:08 +02:00
parent 73d68ab786
commit 8f5f22fe4d

View File

@@ -154,6 +154,31 @@ pub struct PerceptionTrainer {
// Mamba2 encoder block + its optimizer
pub mamba2: Mamba2Block,
pub mamba2_adamw: Mamba2AdamW,
/// Phase 2B: SECOND Mamba2 stack. Sits between the (first-stack-output
/// + LN_a) and the existing LN_b → CfC. Same hidden_dim as stack 1
/// (128), in_dim=HIDDEN_DIM (consumes the LN_a output stream).
pub mamba2_l2: Mamba2Block,
pub mamba2_l2_adamw: Mamba2AdamW,
mamba2_l2_fwd_scratch: Mamba2BlockForwardScratch,
mamba2_l2_bwd_scratch: Mamba2BackwardScratch,
mamba2_l2_grads_buffers: Mamba2BackwardGradsBuffers,
// ── 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>,
ln_a_stats_d: CudaSlice<f32>,
/// `[B, K, HIDDEN_DIM]` — LN_a forward output, fed as input to m2.
ln_a_out_d: GpuTensor,
grad_ln_a_gain_per_row_d: CudaSlice<f32>,
grad_ln_a_bias_per_row_d: CudaSlice<f32>,
grad_ln_a_gain_d: CudaSlice<f32>,
grad_ln_a_bias_d: CudaSlice<f32>,
/// `[B, K, HIDDEN_DIM]` — LN_a backward grad-x: gradient w.r.t.
/// m1's `h_enriched_seq` output (consumed by m1.backward).
grad_ln_a_in_d: GpuTensor,
pub opt_ln_a_gain: AdamW,
pub opt_ln_a_bias: AdamW,
/// Pre-allocated forward intermediates for Mamba2 (input projection
/// output x, A/B projections, h_s2 residual, h_enriched_seq scan
/// output). Construct once at trainer init; reused every step.
@@ -498,6 +523,38 @@ impl PerceptionTrainer {
let mamba2_grads_buffers = Mamba2BackwardGradsBuffers::new(
&stream, cfg.n_batch, cfg.seq_len, FEATURE_DIM, HIDDEN_DIM, cfg.mamba2_state_dim,
).context("Mamba2BackwardGradsBuffers::new")?;
// ── Phase 2B: SECOND Mamba2 stack ──
// Same hidden_dim + state_dim as stack 1. in_dim = HIDDEN_DIM
// (consumes LN_a output which has the same shape as m1's output).
let mamba2_l2 = Mamba2Block::new(
Mamba2BlockConfig {
in_dim: HIDDEN_DIM,
hidden_dim: HIDDEN_DIM,
state_dim: cfg.mamba2_state_dim,
seq_len: cfg.seq_len,
},
stream.clone(),
)
.context("Mamba2Block::new (l2)")?;
let mamba2_l2_adamw = Mamba2AdamW::new(
&mamba2_l2,
Mamba2AdamWConfig {
lr: cfg.lr_mamba2,
grad_clip_max_norm: Some(1.0),
..Default::default()
},
)
.context("Mamba2AdamW::new (l2)")?;
let mamba2_l2_fwd_scratch = Mamba2BlockForwardScratch::new(
&stream, cfg.n_batch, cfg.seq_len, HIDDEN_DIM, HIDDEN_DIM, cfg.mamba2_state_dim,
).context("Mamba2BlockForwardScratch::new (l2)")?;
let mamba2_l2_bwd_scratch = Mamba2BackwardScratch::new(
&stream, cfg.n_batch, cfg.seq_len, HIDDEN_DIM, cfg.mamba2_state_dim,
).context("Mamba2BackwardScratch::new (l2)")?;
let mamba2_l2_grads_buffers = Mamba2BackwardGradsBuffers::new(
&stream, cfg.n_batch, cfg.seq_len, HIDDEN_DIM, HIDDEN_DIM, cfg.mamba2_state_dim,
).context("Mamba2BackwardGradsBuffers::new (l2)")?;
let window_tensor_d = GpuTensor::zeros(&[cfg.n_batch, cfg.seq_len, FEATURE_DIM], &stream)
.map_err(|e| anyhow::anyhow!("window_tensor_d alloc: {e}"))?;
let h_enriched_seq_t_d = GpuTensor::zeros(&[cfg.seq_len, cfg.n_batch, HIDDEN_DIM], &stream)
@@ -586,6 +643,7 @@ impl PerceptionTrainer {
// LayerNorm parameters: gain initialised to 1.0, bias to 0.0.
// Adam wd=0 — LN params are scale/shift, never penalised.
// 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)?;
@@ -594,6 +652,22 @@ 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)?;
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)?;
opt_ln_a_bias.wd = 0.0;
// LN_a + grad scratch buffers (constructed pre-move-of-stream).
let ln_a_stats_d = stream.alloc_zeros::<f32>(cfg.n_batch * cfg.seq_len * 2)?;
let ln_a_out_d = GpuTensor::zeros(&[cfg.n_batch, cfg.seq_len, HIDDEN_DIM], &stream)
.map_err(|e| anyhow::anyhow!("ln_a_out_d alloc: {e}"))?;
let grad_ln_a_gain_per_row_d = stream.alloc_zeros::<f32>(cfg.n_batch * cfg.seq_len * HIDDEN_DIM)?;
let grad_ln_a_bias_per_row_d = stream.alloc_zeros::<f32>(cfg.n_batch * cfg.seq_len * HIDDEN_DIM)?;
let grad_ln_a_gain_d = stream.alloc_zeros::<f32>(HIDDEN_DIM)?;
let grad_ln_a_bias_d = stream.alloc_zeros::<f32>(HIDDEN_DIM)?;
let grad_ln_a_in_d = GpuTensor::zeros(&[cfg.n_batch, cfg.seq_len, HIDDEN_DIM], &stream)
.map_err(|e| anyhow::anyhow!("grad_ln_a_in_d alloc: {e}"))?;
// ── VSN init (Phase 2D) ──
// W_vsn ∈ [FEATURE_DIM, FEATURE_DIM] init near zero so initial
@@ -755,6 +829,23 @@ impl PerceptionTrainer {
mamba2_fwd_scratch,
mamba2_bwd_scratch,
mamba2_grads_buffers,
mamba2_l2,
mamba2_l2_adamw,
mamba2_l2_fwd_scratch,
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,
grad_ln_a_bias_per_row_d,
grad_ln_a_gain_d,
grad_ln_a_bias_d,
grad_ln_a_in_d,
opt_ln_a_gain,
opt_ln_a_bias,
window_tensor_d,
h_enriched_seq_t_d,
grad_h_enriched_seq_t_d,
@@ -795,9 +886,11 @@ impl PerceptionTrainer {
self.opt_heads_b_skip.lr = lr;
}
/// Set Mamba2 AdamW learning rate.
/// Set Mamba2 AdamW learning rate. Applies to BOTH stack-1 and
/// stack-2 (Phase 2B) — they share the same lr_mamba2 config.
pub fn set_lr_mamba2(&mut self, lr: f32) {
self.mamba2_adamw.config.lr = lr;
self.mamba2_l2_adamw.config.lr = lr;
}
/// One training step on a single sequence — thin wrapper around
@@ -1104,17 +1197,14 @@ impl PerceptionTrainer {
unsafe { launch.launch(cfg_vsn).context("variable_selection_fwd")?; }
}
// ── 2. Mamba2 per-step forward — writes into self.mamba2_fwd_scratch.
// Now consumes `vsn_out_d` (gated features) instead of the
// raw `window_tensor_d`.
// ── 2. Mamba2 stack-1 forward — writes into mamba2_fwd_scratch.
// Consumes vsn_out_d (gated features). in_dim = FEATURE_DIM (40).
self.mamba2
.forward_train_seq_into(&self.vsn_out_d, &mut self.mamba2_fwd_scratch)
.context("mamba2 forward_train_seq_into")?;
.context("mamba2 (l1) forward_train_seq_into")?;
// ── 2a. LayerNorm forward over h_enriched_seq [B*K, H].
// Stabilises the trunk-to-CfC distribution. One block per
// row, block-wide tree-reduce for mean/var. Stats saved
// for backward in self.ln_stats_d.
// ── 2a. LayerNorm A — between m1 and m2 (Phase 2B). Reads
// m1.h_enriched_seq, writes ln_a_out_d.
{
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let cfg_ln = LaunchConfig {
@@ -1125,12 +1215,39 @@ 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(&n_rows_ln)
.arg(self.ln_a_out_d.data_mut())
.arg(&mut self.ln_a_stats_d);
unsafe { launch.launch(cfg_ln).context("layer_norm_fwd (LN_a)")?; }
}
// ── 2b. Mamba2 stack-2 forward — consumes ln_a_out_d. in_dim = HIDDEN_DIM.
self.mamba2_l2
.forward_train_seq_into(&self.ln_a_out_d, &mut self.mamba2_l2_fwd_scratch)
.context("mamba2 (l2) forward_train_seq_into")?;
// ── 2c. LayerNorm B — between m2 and CfC. Stabilises the
// trunk-to-CfC distribution. One block per row,
// block-wide tree-reduce for mean/var. Stats saved for
// backward in ln_stats_d. Reads m2's h_enriched_seq.
{
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let cfg_ln = LaunchConfig {
grid_dim: (n_rows_ln as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
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(&n_rows_ln)
.arg(&mut self.ln_out_d)
.arg(&mut self.ln_stats_d);
unsafe { launch.launch(cfg_ln).context("layer_norm_fwd")?; }
unsafe { launch.launch(cfg_ln).context("layer_norm_fwd (LN_b)")?; }
}
// ── 2b. Transpose LN-normalised h_enriched [B, K, H] → [K, B, H]
@@ -1489,8 +1606,8 @@ impl PerceptionTrainer {
unsafe { launch.launch(cfg_tx).context("transpose grad bwd")?; }
}
// ── 7c. LayerNorm backward. Consumes:
// x = Mamba2's h_enriched_seq [B, K, H]
// ── 7c. LayerNorm B backward (between m2 and CfC). Consumes:
// x = m2.h_enriched_seq [B, K, H]
// gain = self.ln_gain_d [H]
// stats = self.ln_stats_d (from fwd) [B*K, 2]
// grad_y = self.grad_h_enriched_seq_d [B, K, H]
@@ -1510,7 +1627,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.mamba2_l2_fwd_scratch.h_enriched_seq.cuda_data())
.arg(&self.ln_gain_d)
.arg(&self.ln_stats_d)
.arg(self.grad_h_enriched_seq_d.cuda_data())
@@ -1518,7 +1635,7 @@ impl PerceptionTrainer {
.arg(self.grad_ln_in_d.data_mut())
.arg(&mut self.grad_ln_gain_per_row_d)
.arg(&mut self.grad_ln_bias_per_row_d);
unsafe { launch.launch(cfg_ln).context("layer_norm_bwd")?; }
unsafe { launch.launch(cfg_ln).context("layer_norm_bwd (LN_b)")?; }
}
// 7c.i — reduce per-row grad_gain → [H].
{
@@ -1551,23 +1668,96 @@ impl PerceptionTrainer {
unsafe { launch.launch(cfg_red).context("layer_norm_reduce bias")?; }
}
// ── 8. Mamba2 backward — fully pre-allocated path. Writes all
// grads into self.mamba2_grads_buffers; no allocation.
// Now consumes:
// - input = vsn_out_d (gated snap_features fed to Mamba2 fwd)
// - d_h_out = grad_ln_in_d (LN bwd output)
// and writes `mamba2_grads_buffers.d_x_from_in` =
// gradient w.r.t. Mamba2's input = gradient on VSN's
// output, which VSN bwd consumes below.
// ── 8. Mamba2 stack-2 backward. Consumes:
// input = ln_a_out_d (m2's forward input)
// d_h_out = grad_ln_in_d (LN_b bwd output)
// Writes:
// mamba2_l2_grads_buffers.d_x_from_in = grad w.r.t. m2's
// input = grad w.r.t. LN_a output (consumed by LN_a bwd below).
self.mamba2_l2
.backward_from_h_enriched_seq_full_into(
&self.ln_a_out_d,
&self.mamba2_l2_fwd_scratch,
&self.grad_ln_in_d,
&mut self.mamba2_l2_bwd_scratch,
&mut self.mamba2_l2_grads_buffers,
)
.context("mamba2 (l2) backward_from_h_enriched_seq_full_into")?;
// ── 8a. LN_a backward. Consumes:
// x = m1.h_enriched_seq [B, K, H]
// gain = ln_a_gain_d [H]
// stats = ln_a_stats_d (from fwd) [B*K, 2]
// grad_y = m2.d_x_from_in (reshape) [B, K, H] — read flat
// Writes:
// grad_x = grad_ln_a_in_d [B, K, H] (fed to m1.bwd)
// grad_gain/row = grad_ln_a_gain_per_row_d
// grad_bias/row = grad_ln_a_bias_per_row_d
{
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let cfg_ln = LaunchConfig {
grid_dim: (n_rows_ln as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
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.ln_a_stats_d)
.arg(self.mamba2_l2_grads_buffers.d_x_from_in.cuda_data())
.arg(&n_rows_ln)
.arg(self.grad_ln_a_in_d.data_mut())
.arg(&mut self.grad_ln_a_gain_per_row_d)
.arg(&mut self.grad_ln_a_bias_per_row_d);
unsafe { launch.launch(cfg_ln).context("layer_norm_bwd (LN_a)")?; }
}
// 8a.i — reduce per-row grad_a_gain → [H].
{
let cfg_red = LaunchConfig {
grid_dim: (HIDDEN_DIM as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let mut launch = self.stream.launch_builder(&self.ln_reduce_fn);
launch
.arg(&self.grad_ln_a_gain_per_row_d)
.arg(&n_rows_ln)
.arg(&mut self.grad_ln_a_gain_d);
unsafe { launch.launch(cfg_red).context("layer_norm_reduce gain (LN_a)")?; }
}
// 8a.ii — reduce per-row grad_a_bias → [H].
{
let cfg_red = LaunchConfig {
grid_dim: (HIDDEN_DIM as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let mut launch = self.stream.launch_builder(&self.ln_reduce_fn);
launch
.arg(&self.grad_ln_a_bias_per_row_d)
.arg(&n_rows_ln)
.arg(&mut self.grad_ln_a_bias_d);
unsafe { launch.launch(cfg_red).context("layer_norm_reduce bias (LN_a)")?; }
}
// ── 8b. Mamba2 stack-1 backward. Consumes:
// input = vsn_out_d (m1's forward input)
// d_h_out = grad_ln_a_in_d (LN_a bwd output)
// Writes:
// mamba2_grads_buffers.d_x_from_in = grad on VSN output
// (consumed by VSN bwd below).
self.mamba2
.backward_from_h_enriched_seq_full_into(
&self.vsn_out_d,
&self.mamba2_fwd_scratch,
&self.grad_ln_in_d,
&self.grad_ln_a_in_d,
&mut self.mamba2_bwd_scratch,
&mut self.mamba2_grads_buffers,
)
.context("mamba2 backward_from_h_enriched_seq_full_into")?;
.context("mamba2 (l1) backward_from_h_enriched_seq_full_into")?;
// ── 8b. VSN backward. Consumes:
// grad_y = mamba2_grads_buffers.d_x_from_in [B*K, FEATURE_DIM]
@@ -1621,13 +1811,18 @@ 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_vsn_w.step(&mut self.vsn_w_d, &self.grad_vsn_w_d)?;
self.opt_vsn_b.step(&mut self.vsn_b_d, &self.grad_vsn_b_d)?;
// GPU-resident grad-clip + AdamW (zero host roundtrips).
// Replaces step_from_buffers which did 9× memcpy_dtoh per step.
self.mamba2_adamw
.step_from_buffers_gpu_clip(&mut self.mamba2, &self.mamba2_grads_buffers)
.context("mamba2 AdamW step_from_buffers_gpu_clip")?;
.context("mamba2 (l1) AdamW step_from_buffers_gpu_clip")?;
self.mamba2_l2_adamw
.step_from_buffers_gpu_clip(&mut self.mamba2_l2, &self.mamba2_l2_grads_buffers)
.context("mamba2 (l2) AdamW step_from_buffers_gpu_clip")?;
// Queue loss_d → loss_host_d (mapped-pinned). Captured inside
// graph; sync + read happen in step_batched outside the region.
@@ -1799,14 +1994,12 @@ impl PerceptionTrainer {
unsafe { launch.launch(cfg_vsn).context("eval variable_selection_fwd")?; }
}
// Mamba2 fwd into pre-allocated fwd_scratch — reads vsn_out_d.
// Mamba2 stack-1 fwd → m1.h_enriched_seq.
self.mamba2
.forward_train_seq_into(&self.vsn_out_d, &mut self.mamba2_fwd_scratch)
.context("eval mamba2 fwd_into")?;
.context("eval mamba2 (l1) fwd_into")?;
// LN forward over h_enriched_seq [B*K, H] — same as training path
// so eval reads the same distribution downstream layers were
// trained on.
// LN_a forward → ln_a_out_d.
{
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let cfg_ln = LaunchConfig {
@@ -1817,12 +2010,38 @@ 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(&n_rows_ln)
.arg(self.ln_a_out_d.data_mut())
.arg(&mut self.ln_a_stats_d);
unsafe { launch.launch(cfg_ln).context("eval layer_norm_fwd (LN_a)")?; }
}
// Mamba2 stack-2 fwd → m2.h_enriched_seq.
self.mamba2_l2
.forward_train_seq_into(&self.ln_a_out_d, &mut self.mamba2_l2_fwd_scratch)
.context("eval mamba2 (l2) fwd_into")?;
// LN_b forward over m2.h_enriched_seq — same as training path so
// eval reads the same distribution downstream layers were
// trained on.
{
let n_rows_ln: i32 = (b_sz * k_seq) as i32;
let cfg_ln = LaunchConfig {
grid_dim: (n_rows_ln as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
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(&n_rows_ln)
.arg(&mut self.ln_out_d)
.arg(&mut self.ln_stats_d);
unsafe { launch.launch(cfg_ln).context("eval layer_norm_fwd")?; }
unsafe { launch.launch(cfg_ln).context("eval layer_norm_fwd (LN_b)")?; }
}
// Transpose LN-normalised h_enriched [B, K, H] → [K, B, H].