From 8f5f22fe4db0d99cad616fab18ef375ab7b8547d Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 17 May 2026 22:28:08 +0200 Subject: [PATCH] feat(ml-alpha): 2-stack Mamba2 with inter-stack LayerNorm (Phase 2B) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- crates/ml-alpha/src/trainer/perception.rs | 281 +++++++++++++++++++--- 1 file changed, 250 insertions(+), 31 deletions(-) diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index b046da706..138c3fce3 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -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, + ln_a_bias_d: CudaSlice, + ln_a_stats_d: CudaSlice, + /// `[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, + grad_ln_a_bias_per_row_d: CudaSlice, + grad_ln_a_gain_d: CudaSlice, + grad_ln_a_bias_d: CudaSlice, + /// `[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 = vec![1.0; HIDDEN_DIM]; let ln_bias_init: Vec = 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::(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::(cfg.n_batch * cfg.seq_len * HIDDEN_DIM)?; + let grad_ln_a_bias_per_row_d = stream.alloc_zeros::(cfg.n_batch * cfg.seq_len * HIDDEN_DIM)?; + let grad_ln_a_gain_d = stream.alloc_zeros::(HIDDEN_DIM)?; + let grad_ln_a_bias_d = stream.alloc_zeros::(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].