diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index fb91c1bcc..c81c78d6f 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -635,12 +635,10 @@ impl IntegratedTrainer { // ── Step 10: encoder grad combine (Phase E.2-DEFER item 1) ─── // Zero the combined slot, fold Q+π+V grad_h_t with λ-weights. - // The combined grad slot is ready for an upcoming follow-up - // commit that lands `PerceptionTrainer::backward_encoder_with_grad_h_t` - // — the encoder-side entry point is a multi-thousand-line refactor - // of `dispatch_train_step`'s backward portion and is being - // sequenced separately to keep this commit focused on the kernel - // signature changes (items 2, 3, 4) that unlock it. + // The combined grad slot feeds `backward_encoder_with_grad_h_t` + // below, which runs the shared encoder backward chain that + // `dispatch_train_step` also uses (single source of truth per + // feedback_single_source_of_truth_no_duplicates). self.stream .memset_zeros(&mut self.grad_h_t_combined_d) .map_err(|e| anyhow::anyhow!("zero grad_h_t_combined_d: {e}"))?; @@ -666,6 +664,21 @@ impl IntegratedTrainer { &mut self.grad_h_t_combined_d, )?; + // ── Step 11: encoder backward (Phase E.3a) ─────────────────── + // Seed the perception trainer's per-K hidden-state grad buffer + // from `grad_h_t_combined_d` at slot K-1 (the only slot the + // RL heads consumed `h_t` from) and run the shared encoder + // backward kernel chain. This updates ALL encoder params: + // CfC ×4, LN ×2, VSN, attn_q, Mamba2 L1+L2 — closes the last + // deferred item from Phase E.2. + self.perception + .backward_encoder_with_grad_h_t( + &self.grad_h_t_combined_d, + self.cfg.perception.n_batch, + self.cfg.perception.seq_len, + ) + .context("perception.backward_encoder_with_grad_h_t")?; + // ── Step 12: compose stats ─────────────────────────────────── // BCE / aux losses are NOT read this phase — perception is driven // separately by callers via its existing step_batched path. The diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index 9eee44f74..3582c7a17 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -589,9 +589,21 @@ pub struct PerceptionTrainer { /// grad_h_old here, which heads_bwd of the previous position adds /// in via the new `grad_h_carry` kernel arg. grad_h_carry_d: CudaSlice, - /// Single-position grad_h_new buffer fed to cfc_step_bwd (heads_bwd - /// writes to it, fused with the carry). - grad_h_new_d: CudaSlice, + /// Phase E.3a: per-K hidden-state gradient accumulator. Replaces the + /// pre-E.3a single-slot `grad_h_new_d` ping-pong with an explicit + /// `[K, B, HIDDEN_DIM]` buffer that head backward kernels write into + /// during Loop 1 (per-K head contributions, NO recurrent carry), and + /// the encoder backward consumes from during Loop 2 (CfC step backward, + /// folding in the recurrent `grad_h_carry_d` per slot before the + /// kernel launch). + /// + /// Makes the encoder backward's input contract explicit so the + /// integrated RL trainer's + /// [`PerceptionTrainer::backward_encoder_with_grad_h_t`] can seed + /// slot `K-1` directly from the caller-provided per-head accumulated + /// `grad_h_t_combined_d` without going through the supervised + /// GRN/aux backward path. The supervised path seeds via Loop 1. + grad_h_new_per_k_d: CudaSlice, /// Single-position zero h_old buffer (used only at k=0 — first /// position of the sequence has no preceding state). zero_h_d: CudaSlice, @@ -2185,7 +2197,11 @@ impl PerceptionTrainer { _horizon_lambda_module: horizon_lambda_module, valid_d: stream.alloc_zeros::(1)?, grad_h_carry_d: stream.alloc_zeros::(cfg.n_batch * n_hid)?, - grad_h_new_d: stream.alloc_zeros::(cfg.n_batch * n_hid)?, + // Phase E.3a: per-K hidden-state gradient buffer + // `[K, B, HIDDEN_DIM]`. Same K × B sizing pattern as + // `h_new_per_k_d` above (per cfg.seq_len = K, cfg.n_batch = B). + // Replaces the pre-E.3a single-slot `grad_h_new_d` buffer. + grad_h_new_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * n_hid)?, zero_h_d: stream.alloc_zeros::(cfg.n_batch * n_hid)?, grad_w_in_d: stream.alloc_zeros::(n_hid * n_in)?, grad_w_rec_d: stream.alloc_zeros::(n_hid * n_hid)?, @@ -3848,6 +3864,588 @@ impl PerceptionTrainer { &self.h_t_d } + /// Phase E.3a (integrated RL trainer): run the encoder backward + /// consuming a caller-provided `grad_h_t` `[B × HIDDEN_DIM]` at slot + /// `K-1`. + /// + /// The integrated trainer + /// ([`crate::trainer::integrated::IntegratedTrainer::step_synthetic`]) + /// calls this AFTER it has run `forward_encoder` + each per-head + /// forward+backward + `grad_h_accumulate` (combining Q+π+V+BCE+aux + /// contributions into a single buffer with loss-balance λ scaling). + /// This method's only job is to seed `grad_h_new_per_k_d` from the + /// caller-provided buffer and then run `dispatch_encoder_backward`. + /// + /// The supervised path (`step_batched` → `dispatch_train_step`) does + /// NOT call this — it seeds `grad_h_new_per_k_d` internally via the + /// GRN head backward kernels (Loop 1) and then calls the same shared + /// `dispatch_encoder_backward` helper (Loop 2). The encoder backward + /// kernel chain is shared between supervised and RL paths per + /// `feedback_single_source_of_truth_no_duplicates`. + /// + /// Seeding semantics: + /// 1. Zero the entire `grad_h_new_per_k_d` buffer (no Loop 1). + /// 2. DtoD-copy `grad_h_t_d` into slot `K-1`. The heads only + /// consumed `h_t = h_new_per_k_d[K-1]`, so this is the only + /// slot where the RL grad lands. All other slots stay zero + /// — those positions had no downstream loss signal. + /// 3. Dispatch the shared encoder backward (Loop 2), which folds + /// the recurrent CfC carry through the K-loop, runs the LN / + /// Mamba2 / VSN / attention chain, reduces, and applies all + /// encoder Adam updates. + pub fn backward_encoder_with_grad_h_t( + &mut self, + grad_h_t_d: &CudaSlice, + b_sz: usize, + k_seq: usize, + ) -> Result<()> { + // 1. Zero the per-K buffer (RL path has no Loop-1 head contributions). + self.stream + .memset_zeros(&mut self.grad_h_new_per_k_d) + .map_err(|e| anyhow::anyhow!("zero grad_h_new_per_k_d: {e}"))?; + + // 2. DtoD-copy grad_h_t_d into slot K-1. Layout in + // grad_h_new_per_k_d is [K, B, HIDDEN_DIM] row-major (k + // outermost), matching `h_new_per_k_d`'s layout (see field + // docs at line ~362). Slot K-1 lives at offset + // (K-1) * B * HIDDEN_DIM floats. + let hidden_dim = HIDDEN_DIM; + let nbytes = b_sz * hidden_dim * std::mem::size_of::(); + let dst_offset_bytes = + (k_seq - 1) * b_sz * hidden_dim * std::mem::size_of::(); + unsafe { + let (dst_base_ptr, _g_dst) = + self.grad_h_new_per_k_d.device_ptr_mut(&self.stream); + let (src_ptr, _g_src) = grad_h_t_d.device_ptr(&self.stream); + let dst_ptr = dst_base_ptr + dst_offset_bytes as u64; + cudarc::driver::result::memcpy_dtod_async( + dst_ptr, + src_ptr, + nbytes, + self.stream.cu_stream(), + ) + .map_err(|e| { + anyhow::anyhow!("memcpy_dtod grad_h_t into slot K-1: {e}") + })?; + } + + // 3. Run the shared Loop 2 (encoder backward). + self.dispatch_encoder_backward(b_sz, k_seq)?; + Ok(()) + } + + /// Phase E.3a: encoder-only backward pass — Loop 2 of the split + /// supervised K-loop and the sole encoder-backward entry point for + /// the RL path. + /// + /// Reads from `grad_h_new_per_k_d` (per-K accumulator seeded by + /// Loop 1 in the supervised path, or by + /// [`Self::backward_encoder_with_grad_h_t`] in the RL path) and + /// runs the full encoder backward kernel chain: + /// - K-loop CfC step backward (slot k of `grad_h_new_per_k_d`, + /// folds in `grad_h_carry_d` recurrent carry before kernel + /// launch; writes encoder param grad scratches + `grad_h_carry` + /// for slot k-1 + `grad_h_enriched_seq_t_d[k]`). + /// - Transpose `grad_h_enriched_seq_t_d` `[K, B, H]` → `[B, K, H]`. + /// - Attention-pool backward + reducer (`grad_attn_q_d`). + /// - LN_b backward + 2 reducers (`grad_ln_gain_d`, + /// `grad_ln_bias_d`). + /// - Mamba2 stack-2 backward (encoder grads buffer). + /// - LN_a backward + 2 reducers. + /// - Mamba2 stack-1 backward. + /// - VSN backward + 2 reducers (`grad_vsn_w_d`, `grad_vsn_b_d`). + /// - `reduce_axis0` for encoder param grads (CfC w_in/w_rec/b/tau). + /// - Encoder Adam steps: CfC ×4, Controller B (τ projection), + /// LN_b gain/bias, LN_a gain/bias, VSN w/b, attn_q, + /// Mamba2 L1 + L2 grouped GPU-resident Adam. + /// + /// Does NOT touch: + /// - GRN head backward (Loop 1's job, in `dispatch_train_step`). + /// - Aux head + trunk backward (separate K-loop in + /// `dispatch_train_step`). + /// - Smoothness backward (upstream, before Loop 1). + /// - BCE/GRN head Adam, Controllers C, ALPHA fix — all wrap the + /// BCE head Adam in `dispatch_train_step` after this helper + /// returns. + fn dispatch_encoder_backward( + &mut self, + b_sz: usize, + k_seq: usize, + ) -> Result<()> { + // ── Re-derive per-K slot byte strides used across this helper. + let kb_hid_bytes = b_sz * HIDDEN_DIM * std::mem::size_of::(); + let n_in_i = HIDDEN_DIM as i32; + let n_hid_i = HIDDEN_DIM as i32; + let n_batch_i = b_sz as i32; + let dt_s = 1.0_f32; + + // CfC bwd launch config (block-per-batch + cfc_bwd_smem shared mem). + // Shared layout: sd_pre[n_hid] + sdecay[n_hid] + s_x[n_in] + s_h_old[n_hid] + // = 4 × HIDDEN_DIM floats per block. + let cfc_bwd_smem = (4 * HIDDEN_DIM * std::mem::size_of::()) as u32; + let cfg_cfc_bwd = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (HIDDEN_DIM as u32, 1, 1), + shared_mem_bytes: cfc_bwd_smem, + }; + + // ── 6.L2 Loop 2: reverse-order K loop running cfc_step_bwd ───── + // at each k. Folds `grad_h_carry_d` (recurrent carry from + // k+1's cfc bwd) into `grad_h_new_per_k_d[k]` BEFORE the + // kernel launch — preserves the prior fused-loop semantics + // byte-identically: the cfc_step_bwd kernel sees exactly + // `head_contrib[k] + carry_from_k+1`, just like before + // (under the old fused path the GRN bwd kernel's + // grad_h_carry arg did this fold). + // + // cfc_step_bwd_batched writes: + // - grad_h_carry_d (= grad_h_old_out for position k-1) + // - grad_x[B, n_in] → directly into grad_h_enriched_seq_t_d[k] + // - per-batch param grad scratches via += accumulation + // Caller (step_batched / step_from_buffers_gpu_clip) zeroes + // the param scratches once per step. + self.stream + .memset_zeros(&mut self.grad_h_carry_d) + .map_err(|e| anyhow::anyhow!("zero grad_h_carry: {e}"))?; + + let attn_context_ptr_bwd = { + let (p, _g) = self.attn_context_d.device_ptr(&self.stream); + p + }; + let henr_t_base_bwd = { + let (p, _g) = self.h_enriched_seq_t_d.cuda_data().device_ptr(&self.stream); + p + }; + let (h_per_k_base_bwd, _g_hpk_bwd) = + self.h_new_per_k_d.device_ptr_mut(&self.stream); + let (grad_henr_t_base, _g_ghen_t_mut) = + self.grad_h_enriched_seq_t_d.data_mut().device_ptr_mut(&self.stream); + let (grad_h_new_per_k_base, _g_ghnk) = + self.grad_h_new_per_k_d.device_ptr_mut(&self.stream); + + let n_carry_elems = (b_sz * HIDDEN_DIM) as i32; + let carry_block: u32 = 256; + let carry_grid: u32 = ((n_carry_elems as u32) + carry_block - 1) / carry_block; + let cfg_carry_add = LaunchConfig { + grid_dim: (carry_grid, 1, 1), + block_dim: (carry_block, 1, 1), + shared_mem_bytes: 0, + }; + + for k in (0..k_seq).rev() { + let x_k_ptr = henr_t_base_bwd + (k * kb_hid_bytes) as u64; + let h_old_k_ptr = if k == 0 { + attn_context_ptr_bwd + } else { + h_per_k_base_bwd + ((k - 1) * kb_hid_bytes) as u64 + }; + let grad_henr_k_ptr = grad_henr_t_base + (k * kb_hid_bytes) as u64; + // Slot k of the per-K grad accumulator — cfc_step_bwd reads + // this as `grad_h_new`. Fold the recurrent carry into it + // first via aux_vec_add_inplace (dst += src). + let grad_h_new_k_ptr = grad_h_new_per_k_base + (k * kb_hid_bytes) as u64; + + // Fold grad_h_carry_d INTO grad_h_new_per_k_d[k]. At k=K-1, + // grad_h_carry_d is zero (we just memset it above), so this + // is a no-op for the first iteration. For k < K-1, + // grad_h_carry_d holds the carry from k+1's cfc_step_bwd. + // Byte-identical to the prior fused-loop semantics where + // the GRN bwd kernel folded grad_h_carry into its grad_h + // output (see line 814 of multi_horizon_heads.cu). + unsafe { + let mut launch = self.stream.launch_builder(&self.aux_vec_add_fn); + launch + .arg(&grad_h_new_k_ptr) + .arg(&self.grad_h_carry_d) + .arg(&n_carry_elems); + launch + .launch(cfg_carry_add) + .context("Loop 2: grad_h_new_per_k[k] += grad_h_carry")?; + } + + // cfc_step_bwd_batched: reads grad_h_new_per_k_d[k] (with + // carry folded in above), writes grad_h_carry_d (= grad_h_old + // for k-1) and grad_x (= grad_henr_k_ptr, the K-slot of + // grad_h_enriched_seq_t_d). Per-batch param-grad scratch + // accumulates via += across K iterations; reduce_axis0 + // collapses B → final grad after the K-loop. + unsafe { + let mut launch = + self.stream.launch_builder(&self.step_bwd_batched_fn); + launch + .arg(&self.trunk.w_in_d) + .arg(&self.trunk.w_rec_d) + .arg(&self.trunk.b_d) + .arg(&self.trunk.tau_all_d) + .arg(&x_k_ptr) + .arg(&h_old_k_ptr) + .arg(&grad_h_new_k_ptr) + .arg(&dt_s) + .arg(&n_in_i) + .arg(&n_hid_i) + .arg(&n_batch_i) + .arg(&mut self.cfc_grad_w_in_scratch_d) + .arg(&mut self.cfc_grad_w_rec_scratch_d) + .arg(&mut self.cfc_grad_b_scratch_d) + .arg(&mut self.cfc_grad_tau_scratch_d) + .arg(&mut self.grad_h_carry_d) + .arg(&grad_henr_k_ptr); + launch + .launch(cfg_cfc_bwd) + .context("Loop 2: cfc_step_bwd k")?; + } + } + drop((_g_hpk_bwd, _g_ghen_t_mut, _g_ghnk)); + + // ── 7b. Transpose grad_h_enriched_seq_t_d [K, B, H] → [B, K, H] + // (pre-allocated grad_h_enriched_seq_d). This is the + // gradient w.r.t. the LN OUTPUT (= grad_y for LN bwd). + { + let block_n3: u32 = 32; + let grid_z = (HIDDEN_DIM as u32).div_ceil(block_n3); + let cfg_tx = LaunchConfig { + grid_dim: (k_seq as u32, b_sz as u32, grid_z), + block_dim: (block_n3, 1, 1), + shared_mem_bytes: 0, + }; + let n1 = k_seq as i32; + let n2 = b_sz as i32; + let n3 = HIDDEN_DIM as i32; + let mut launch = self.stream.launch_builder(&self.trunk.transpose_3d_fn); + launch + .arg(self.grad_h_enriched_seq_t_d.cuda_data()) + .arg(self.grad_h_enriched_seq_d.data_mut()) + .arg(&n1).arg(&n2).arg(&n3); + unsafe { launch.launch(cfg_tx).context("transpose grad bwd")?; } + } + + // ── 7c-pre. Attention pool backward (Phase 3). Consumes: + // Q = self.attn_q_d [HIDDEN_DIM] + // ln_out (values)= self.ln_out_d [B, K, HIDDEN_DIM] + // attn_weights = self.attn_weights_d (saved by fwd) [B, K] + // grad_context = self.grad_h_carry_d (= grad on initial + // h_old at k=0, which IS attn_context) [B, HIDDEN_DIM] + // Writes (BOTH ARE +=): + // grad_attn_q_d += attn-path contribution to Q + // grad_h_enriched_seq_d (LN_b output grad) += attn-path + // contribution to ln_out + // The pre-zero of grad_attn_q_d at step start makes the += a + // clean overwrite for the Q grad. grad_h_enriched_seq_d already + // holds the K-loop's contribution at this point — attn's + // contribution adds on top. + // Phase B commit 4: block-per-batch attn bwd writes per-batch + // grad_Q scratch; reducer collapses → final grad_attn_q_d below. + { + let k_i32 = k_seq as i32; + let n_batch_attn = b_sz as i32; + let shared = (2 * k_seq + 128) * std::mem::size_of::(); + let cfg_attn_bwd = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (128, 1, 1), // ATTN_BLOCK + shared_mem_bytes: shared as u32, + }; + let mut launch = self.stream.launch_builder(&self.attn_bwd_fn); + launch + .arg(&self.trunk.attn_q_d) + .arg(&self.ln_out_d) + .arg(&self.attn_weights_d) + .arg(&self.grad_h_carry_d) + .arg(&n_batch_attn).arg(&k_i32) + .arg(&mut self.attn_grad_q_scratch_d) + .arg(self.grad_h_enriched_seq_d.data_mut()); + unsafe { launch.launch(cfg_attn_bwd).context("attention_pool_bwd")?; } + } + // Attn pool reducer: collapse [B, HIDDEN_DIM] → [HIDDEN_DIM]. + { + let n_tail_i = HIDDEN_DIM as i32; + let cfg_red = LaunchConfig { + grid_dim: (((HIDDEN_DIM + 31) / 32) as u32, 1, 1), + block_dim: (32, 8, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); + launch + .arg(&self.attn_grad_q_scratch_d) + .arg(&n_batch_i) + .arg(&n_tail_i) + .arg(&mut self.grad_attn_q_d); + unsafe { launch.launch(cfg_red).context("reduce attn_grad_q")?; } + } + + // ── 7c. LayerNorm B backward (between m2 and CfC). + { + 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_l2_fwd_scratch.h_enriched_seq.cuda_data()) + .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) + .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 (LN_b)")?; } + } + // 7c.i — reduce per-row grad_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_gain_per_row_d) + .arg(&n_rows_ln) + .arg(&mut self.grad_ln_gain_d); + unsafe { launch.launch(cfg_red).context("layer_norm_reduce gain")?; } + } + // 7c.ii — reduce per-row grad_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_bias_per_row_d) + .arg(&n_rows_ln) + .arg(&mut self.grad_ln_bias_d); + unsafe { launch.launch(cfg_red).context("layer_norm_reduce bias")?; } + } + + // ── 8. Mamba2 stack-2 backward. + self.trunk + .mamba2_l2_mut() + .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. + { + 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.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) + .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. + self.trunk + .mamba2_l1_mut() + .backward_from_h_enriched_seq_full_into( + &self.vsn_out_d, + &self.mamba2_fwd_scratch, + &self.grad_ln_a_in_d, + &mut self.mamba2_bwd_scratch, + &mut self.mamba2_grads_buffers, + ) + .context("mamba2 (l1) backward_from_h_enriched_seq_full_into")?; + + // ── 8b. VSN backward. + self.stream.memset_zeros(&mut self.vsn_grad_w_scratch_d) + .map_err(|e| anyhow::anyhow!("zero vsn_grad_w_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.vsn_grad_b_scratch_d) + .map_err(|e| anyhow::anyhow!("zero vsn_grad_b_scratch: {e}"))?; + { + let n_rows_vsn: i32 = (b_sz * k_seq) as i32; + let cfg_vsn_bwd = LaunchConfig { + grid_dim: (n_rows_vsn as u32, 1, 1), + block_dim: (64, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.vsn_bwd_fn); + launch + .arg(&self.trunk.vsn_w_d) + .arg(self.window_tensor_d.cuda_data()) + .arg(&self.vsn_gates_d) + .arg(self.mamba2_grads_buffers.d_x_from_in.cuda_data()) + .arg(&n_rows_vsn) + .arg(&mut self.vsn_grad_w_scratch_d) + .arg(&mut self.vsn_grad_b_scratch_d) + .arg(&mut self.vsn_grad_x_d); + unsafe { launch.launch(cfg_vsn_bwd).context("variable_selection_bwd")?; } + } + // VSN reducer: collapse n_rows (= B * K) → final grad buffers. + { + let n_rows_i = (b_sz * k_seq) as i32; + let n_tail_w_usz = FEATURE_DIM * FEATURE_DIM; + let cfg_red_w = LaunchConfig { + grid_dim: (((n_tail_w_usz + 31) / 32) as u32, 1, 1), + block_dim: (32, 8, 1), + shared_mem_bytes: 0, + }; + let n_tail_w = n_tail_w_usz as i32; + unsafe { + let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); + launch + .arg(&self.vsn_grad_w_scratch_d) + .arg(&n_rows_i) + .arg(&n_tail_w) + .arg(&mut self.grad_vsn_w_d); + launch.launch(cfg_red_w).context("reduce vsn_grad_w")?; + } + let cfg_red_b = LaunchConfig { + grid_dim: (((FEATURE_DIM + 31) / 32) as u32, 1, 1), + block_dim: (32, 8, 1), + shared_mem_bytes: 0, + }; + let n_tail_b = FEATURE_DIM as i32; + unsafe { + let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); + launch + .arg(&self.vsn_grad_b_scratch_d) + .arg(&n_rows_i) + .arg(&n_tail_b) + .arg(&mut self.grad_vsn_b_d); + launch.launch(cfg_red_b).context("reduce vsn_grad_b")?; + } + } + + // ── 8c. (Phase B) Reduce CfC per-batch grad scratch → final grad buffers. + // 4 launches, one per cfc param tensor. Phase E.3a: this is + // the ENCODER half of the prior combined reducer closure; + // the HEAD half (GRN ×10) stays inline in dispatch_train_step. + { + let reduce_encoder = |n_tail: usize, + scratch: &CudaSlice, + out: &mut CudaSlice, + label: &'static str| + -> Result<()> { + let cfg = LaunchConfig { + grid_dim: (((n_tail + 31) / 32) as u32, 1, 1), + block_dim: (32, 8, 1), + shared_mem_bytes: 0, + }; + let n_tail_i = n_tail as i32; + let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); + launch + .arg(scratch) + .arg(&n_batch_i) + .arg(&n_tail_i) + .arg(out); + unsafe { launch.launch(cfg).context(label)?; } + Ok(()) + }; + reduce_encoder(HIDDEN_DIM * HIDDEN_DIM, &self.cfc_grad_w_in_scratch_d, + &mut self.grad_w_in_d, "reduce cfc_grad_w_in")?; + reduce_encoder(HIDDEN_DIM * HIDDEN_DIM, &self.cfc_grad_w_rec_scratch_d, + &mut self.grad_w_rec_d, "reduce cfc_grad_w_rec")?; + reduce_encoder(HIDDEN_DIM, &self.cfc_grad_b_scratch_d, + &mut self.grad_b_d, "reduce cfc_grad_b")?; + reduce_encoder(HIDDEN_DIM, &self.cfc_grad_tau_scratch_d, + &mut self.grad_tau_d, "reduce cfc_grad_tau")?; + } + + // ── 9. Apply ENCODER AdamW updates: CfC ×4 + Controller B + // (τ projection) + LN_b gain/bias + LN_a gain/bias + + // VSN w/b + attn_q + Mamba2 L1+L2 grouped Adam. + self.opt_w_in.step(&mut self.trunk.w_in_d, &self.grad_w_in_d)?; + self.opt_w_rec.step(&mut self.trunk.w_rec_d, &self.grad_w_rec_d)?; + self.opt_b.step(&mut self.trunk.b_d, &self.grad_b_d)?; + self.opt_tau.step(&mut self.trunk.tau_all_d, &self.grad_tau_d)?; + // ── Controller B: post-Adam τ projection (spec §3.2) ── + // Per `pearl_adam_normalizes_loss_weights`, the only way to actually + // bound `tau_all_d` is a hard projection AFTER Adam updates the + // parameter. Loss-weight modulation would no-op because Adam's + // m/sqrt(v) cancels constant multipliers. Phase 2 only — Phase 1's + // captured graph never enters this branch since the transition + // invalidates the graph (`pearl_no_host_branches_in_captured_graph`). + if self.phase == TrainingPhase::Phase2Routed { + if let Some(metadata) = self.bucket_routing_metadata.as_ref() { + let cfg_tau_clamp = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (HIDDEN_DIM as u32, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.tau_clamp_fn); + launch + .arg(&mut self.trunk.tau_all_d) + .arg(&metadata.bucket_id_per_channel_d) + .arg(&metadata.bucket_tau_iqr_lo_d) + .arg(&metadata.bucket_tau_iqr_hi_d); + unsafe { + launch + .launch(cfg_tau_clamp) + .context("tau_clamp_kernel launch (Controller B)")?; + } + } + } + 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)?; + self.opt_vsn_b.step(&mut self.trunk.vsn_b_d, &self.grad_vsn_b_d)?; + self.opt_attn_q.step(&mut self.trunk.attn_q_d, &self.grad_attn_q_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(self.trunk.mamba2_l1_mut(), &self.mamba2_grads_buffers) + .context("mamba2 (l1) AdamW step_from_buffers_gpu_clip")?; + self.mamba2_l2_adamw + .step_from_buffers_gpu_clip(self.trunk.mamba2_l2_mut(), &self.mamba2_l2_grads_buffers) + .context("mamba2 (l2) AdamW step_from_buffers_gpu_clip")?; + Ok(()) + } + /// Kernel-dispatch portion of a training step — captured into the /// CUDA Graph on the second `step_batched` call. fn dispatch_train_step( @@ -4237,16 +4835,9 @@ impl PerceptionTrainer { block_dim: (N_HORIZONS as u32, 1, 1), shared_mem_bytes: 0, }; - // CfC bwd (Phase B): block-per-batch + 2 * n_hid floats of shared mem - // (one sd_pre row + sdecay for the block's single bi). - // Shared layout: sd_pre [n_hid] + sdecay [n_hid] + s_x [n_in] + s_h_old [n_hid]. - // With n_in == n_hid == HIDDEN_DIM, that's 4 × HIDDEN_DIM floats per block. - let cfc_bwd_smem = (4 * HIDDEN_DIM * std::mem::size_of::()) as u32; - let cfg_cfc_bwd = LaunchConfig { - grid_dim: (b_sz as u32, 1, 1), - block_dim: (HIDDEN_DIM as u32, 1, 1), - shared_mem_bytes: cfc_bwd_smem, - }; + // Phase E.3a: `cfc_bwd_smem` + `cfg_cfc_bwd` moved into + // `dispatch_encoder_backward` (Loop 2) where the cfc_step_bwd + // K-loop lives now. // Batched heads backward shared mem: sd_z[B, 5]. let heads_bwd_smem = (b_sz * N_HORIZONS * std::mem::size_of::()) as u32; let cfg_heads_bwd = LaunchConfig { @@ -4834,39 +5425,44 @@ impl PerceptionTrainer { unsafe { launch.launch(cfg).context("horizon_ema_and_lambda launch")?; } } - // ── 6. Reverse-order backward K loop using pre-allocated - // grad_h_enriched_seq_t_d as the per-K slot output. - self.stream.memset_zeros(&mut self.grad_h_carry_d) - .map_err(|e| anyhow::anyhow!("zero grad_h_carry: {e}"))?; + // ── 6. Loop 1 (head backwards) — per-K head contributions ────── + // Phase E.3a refactor: split the previously-fused + // (GRN bwd) → (CfC bwd) reverse K-loop into two passes. + // + // Loop 1 writes the pure per-K head gradient w.r.t. the CfC + // hidden state into `grad_h_new_per_k_d[k]`. The GRN bwd + // kernel is called with `grad_h_carry = nullptr` so the + // output is the head contribution alone — NO recurrent + // carry from k+1's CfC bwd. The carry is folded back in + // during Loop 2 (in `dispatch_encoder_backward`) where + // cfc_step_bwd is launched — see the helper for the + // slot-wise `aux_vec_add_inplace` that adds + // `grad_h_carry_d` into `grad_h_new_per_k_d[k]` before + // cfc_step_bwd reads it. Byte-identical to the prior + // fused-loop semantics: the final value seen by + // cfc_step_bwd is still `head_contrib[k] + carry_from_k+1`. + // + // Splitting makes the encoder backward's input contract + // explicit (an `[K, B, HIDDEN_DIM]` buffer) so the + // integrated RL trainer's + // `backward_encoder_with_grad_h_t` can seed slot K-1 from + // a caller-provided combined grad without going through + // the supervised GRN bwd at all. + self.stream.memset_zeros(&mut self.grad_h_new_per_k_d) + .map_err(|e| anyhow::anyhow!("zero grad_h_new_per_k_d: {e}"))?; - // Phase 3: at k=0 the bwd kernel reads `h_old` = attn_context_d - // (mirroring the forward pass). zero_h_ptr_bwd retained as a - // legacy fallback / unused alias. - let _zero_h_ptr_bwd_unused = { - let (p, _g) = self.zero_h_d.device_ptr(&self.stream); - p - }; - let attn_context_ptr_bwd = { - let (p, _g) = self.attn_context_d.device_ptr(&self.stream); - p - }; - let henr_t_base_bwd = { - let (p, _g) = self.h_enriched_seq_t_d.cuda_data().device_ptr(&self.stream); - p - }; - let (h_per_k_base_bwd, _g_hpk_bwd) = self.h_new_per_k_d.device_ptr_mut(&self.stream); let (probs_base_bwd, _g_probs_bwd) = self.probs_per_k_d.device_ptr_mut(&self.stream); let (gprobs_base_bwd, _g_gprobs_bwd) = self.grad_probs_per_k_d.device_ptr_mut(&self.stream); - let (grad_henr_t_base, _g_ghen_t_mut) = self.grad_h_enriched_seq_t_d.data_mut().device_ptr_mut(&self.stream); + let (h_per_k_base_bwd_l1, _g_hpk_bwd_l1) = self.h_new_per_k_d.device_ptr_mut(&self.stream); let (z1_base_bwd, _g_z1_bwd) = self.z1_per_k_d.device_ptr_mut(&self.stream); let (a1_base_bwd, _g_a1_bwd) = self.a1_per_k_d.device_ptr_mut(&self.stream); let (z2_base_bwd, _g_z2_bwd) = self.z2_per_k_d.device_ptr_mut(&self.stream); let (gate_base_bwd, _g_gate_bwd) = self.gate_logit_per_k_d.device_ptr_mut(&self.stream); let (main_base_bwd, _g_main_bwd) = self.main_per_k_d.device_ptr_mut(&self.stream); + let (grad_h_new_per_k_base, _g_ghnk) = self.grad_h_new_per_k_d.device_ptr_mut(&self.stream); for k in (0..k_seq).rev() { - let h_new_k_ptr = h_per_k_base_bwd + (k * kb_hid_bytes) as u64; - let x_k_ptr = henr_t_base_bwd + (k * kb_hid_bytes) as u64; + let h_new_k_ptr = h_per_k_base_bwd_l1 + (k * kb_hid_bytes) as u64; let probs_k_ptr = probs_base_bwd + (k * kb_nh_bytes) as u64; let gprobs_k_ptr = gprobs_base_bwd + (k * kb_nh_bytes) as u64; let z1_k_ptr = z1_base_bwd + (k * kb_nh_mid_bytes) as u64; @@ -4874,12 +5470,8 @@ impl PerceptionTrainer { let z2_k_ptr = z2_base_bwd + (k * kb_nh_mid_bytes) as u64; let gate_k_ptr = gate_base_bwd + (k * kb_nh_bytes) as u64; let main_k_ptr = main_base_bwd + (k * kb_nh_bytes) as u64; - let h_old_k_ptr = if k == 0 { - attn_context_ptr_bwd - } else { - h_per_k_base_bwd + ((k - 1) * kb_hid_bytes) as u64 - }; - let grad_henr_k_ptr = grad_henr_t_base + (k * kb_hid_bytes) as u64; + // Output slot for THIS k: grad_h_new_per_k_d[k] → [B, HIDDEN_DIM]. + let grad_h_new_k_ptr = grad_h_new_per_k_base + (k * kb_hid_bytes) as u64; // GRN backward: chain rule through sigmoid → (skip + sigmoid(gate)*main) // → GLU → linear → GELU → linear → trunk. `lambda_d` scales @@ -4889,6 +5481,13 @@ impl PerceptionTrainer { // Phase B: per-batch GRN grad scratch + per-batch grad_h_new. // K-loop's 64 invocations accumulate into the scratch via +=; // reduce_axis0 collapses → final grad after the K-loop. + // Phase E.3a: pass `nullptr` (= 0) for `grad_h_carry` — the + // GRN kernel's null check (line ~814 of multi_horizon_heads.cu) + // treats nullptr as zero carry, so the output `grad_h` slot + // contains the pure per-K head contribution. Loop 2 folds in + // the recurrent carry from k+1's cfc_step_bwd before + // consuming slot k. + let null_carry_ptr: u64 = 0; unsafe { let mut launch = self.stream.launch_builder(&self.heads_grn_bwd_fn); launch @@ -4899,7 +5498,7 @@ impl PerceptionTrainer { .arg(&z1_k_ptr).arg(&a1_k_ptr).arg(&z2_k_ptr) .arg(&gate_k_ptr).arg(&main_k_ptr) .arg(&h_new_k_ptr) - .arg(&self.grad_h_carry_d) + .arg(&null_carry_ptr) .arg(&self.lambda_d) .arg(&n_batch_i) .arg(&mut self.grn_grad_w1_scratch_d).arg(&mut self.grn_grad_b1_scratch_d) @@ -4907,354 +5506,54 @@ impl PerceptionTrainer { .arg(&mut self.grn_grad_w_gate_scratch_d).arg(&mut self.grn_grad_b_gate_scratch_d) .arg(&mut self.grn_grad_w_main_scratch_d).arg(&mut self.grn_grad_b_main_scratch_d) .arg(&mut self.grn_grad_w_skip_scratch_d).arg(&mut self.grn_grad_b_skip_scratch_d) - .arg(&mut self.grad_h_new_d); - launch.launch(cfg_grn_bwd).context("heads GRN bwd k")?; - } - - // cfc_step_bwd_batched: writes grad_h_carry (= grad_h_old_out for - // position k-1) and grad_x[B, n_in]. grad_x output buffer IS - // grad_henr_k_ptr — kernel writes directly into the K-slot. - // Phase B: per-batch param-grad scratch + per-batch grad_h_old / grad_x. - // Writes to scratch via += across K iterations; reduce_axis0 - // collapses B → final grad after the K-loop (see Pass 8c). - unsafe { - let mut launch = self.stream.launch_builder(&self.step_bwd_batched_fn); - launch - .arg(&self.trunk.w_in_d).arg(&self.trunk.w_rec_d).arg(&self.trunk.b_d).arg(&self.trunk.tau_all_d) - .arg(&x_k_ptr).arg(&h_old_k_ptr).arg(&self.grad_h_new_d) - .arg(&dt_s).arg(&n_in_i).arg(&n_hid_i).arg(&n_batch_i) - .arg(&mut self.cfc_grad_w_in_scratch_d) - .arg(&mut self.cfc_grad_w_rec_scratch_d) - .arg(&mut self.cfc_grad_b_scratch_d) - .arg(&mut self.cfc_grad_tau_scratch_d) - .arg(&mut self.grad_h_carry_d) - .arg(&grad_henr_k_ptr); - launch.launch(cfg_cfc_bwd).context("cfc bwd k batched")?; + .arg(&grad_h_new_k_ptr); + launch.launch(cfg_grn_bwd).context("heads GRN bwd k (Loop 1)")?; } } - drop((_g_hpk_bwd, _g_probs_bwd, _g_gprobs_bwd, _g_ghen_t_mut, - _g_z1_bwd, _g_a1_bwd, _g_z2_bwd, _g_gate_bwd, _g_main_bwd)); + drop((_g_probs_bwd, _g_gprobs_bwd, _g_hpk_bwd_l1, + _g_z1_bwd, _g_a1_bwd, _g_z2_bwd, _g_gate_bwd, _g_main_bwd, + _g_ghnk)); - // ── 7. Stream stays async — bwd loop kernels, transposes, and - // Mamba2 bwd are all sequential on the same stream, no - // sync needed between them. Single sync at the very end - // to flush before reading the mapped-pinned loss shadow. + // ── 6b. Loop 2 (encoder backward) ────────────────────────────── + // Single source of truth per + // `feedback_single_source_of_truth_no_duplicates`. The + // integrated RL trainer's `backward_encoder_with_grad_h_t` + // seeds `grad_h_new_per_k_d` differently (slot K-1 only, + // all other slots zero) and then calls the SAME helper. + self.dispatch_encoder_backward(b_sz, k_seq)?; - // ── 7b. Transpose grad_h_enriched_seq_t_d [K, B, H] → [B, K, H] - // (pre-allocated grad_h_enriched_seq_d). This is the - // gradient w.r.t. the LN OUTPUT (= grad_y for LN bwd). - { - let block_n3: u32 = 32; - let grid_z = (HIDDEN_DIM as u32).div_ceil(block_n3); - let cfg_tx = LaunchConfig { - grid_dim: (k_seq as u32, b_sz as u32, grid_z), - block_dim: (block_n3, 1, 1), - shared_mem_bytes: 0, - }; - let n1 = k_seq as i32; - let n2 = b_sz as i32; - let n3 = HIDDEN_DIM as i32; - let mut launch = self.stream.launch_builder(&self.trunk.transpose_3d_fn); - launch - .arg(self.grad_h_enriched_seq_t_d.cuda_data()) - .arg(self.grad_h_enriched_seq_d.data_mut()) - .arg(&n1).arg(&n2).arg(&n3); - unsafe { launch.launch(cfg_tx).context("transpose grad bwd")?; } - } + // ── 7. Stream stays async — head bwd loop, encoder helper, head + // reducers, and head Adams are all sequential on the same + // stream, no sync needed between them. Single sync at the + // very end to flush before reading the mapped-pinned loss + // shadow. + // + // Phase E.3a: sections 7b–8b (transpose, attn pool bwd, + // LN bwd ×2, Mamba2 bwd ×2, VSN bwd) + encoder reducers + // (CfC ×4) + encoder Adams (CfC ×4 + Controller B + + // LN ×2 + VSN + attn_q + Mamba2 L1/L2) have moved into + // `dispatch_encoder_backward` so the supervised and RL + // paths share a single encoder-backward implementation. + // Only the HEAD reducers (GRN ×10) and HEAD Adams + // (GRN ×10 + Controller C + ALPHA fix) remain inline + // below — they live with the head loss bwd K-loop (Loop 1) + // that produces their inputs. - // ── 7c-pre. Attention pool backward (Phase 3). Consumes: - // Q = self.attn_q_d [HIDDEN_DIM] - // ln_out (values)= self.ln_out_d [B, K, HIDDEN_DIM] - // attn_weights = self.attn_weights_d (saved by fwd) [B, K] - // grad_context = self.grad_h_carry_d (= grad on initial - // h_old at k=0, which IS attn_context) [B, HIDDEN_DIM] - // Writes (BOTH ARE +=): - // grad_attn_q_d += attn-path contribution to Q - // grad_h_enriched_seq_d (LN_b output grad) += attn-path - // contribution to ln_out - // The pre-zero of grad_attn_q_d at step start makes the += a - // clean overwrite for the Q grad. grad_h_enriched_seq_d already - // holds the K-loop's contribution at this point — attn's - // contribution adds on top. - // Phase B commit 4: block-per-batch attn bwd writes per-batch - // grad_Q scratch; reducer collapses → final grad_attn_q_d below. - { - let k_i32 = k_seq as i32; - let n_batch_attn = b_sz as i32; - let shared = (2 * k_seq + 128) * std::mem::size_of::(); - let cfg_attn_bwd = LaunchConfig { - grid_dim: (b_sz as u32, 1, 1), - block_dim: (128, 1, 1), // ATTN_BLOCK - shared_mem_bytes: shared as u32, - }; - let mut launch = self.stream.launch_builder(&self.attn_bwd_fn); - launch - .arg(&self.trunk.attn_q_d) - .arg(&self.ln_out_d) - .arg(&self.attn_weights_d) - .arg(&self.grad_h_carry_d) - .arg(&n_batch_attn).arg(&k_i32) - .arg(&mut self.attn_grad_q_scratch_d) - .arg(self.grad_h_enriched_seq_d.data_mut()); - unsafe { launch.launch(cfg_attn_bwd).context("attention_pool_bwd")?; } - } - // Attn pool reducer: collapse [B, HIDDEN_DIM] → [HIDDEN_DIM]. - // reduce_axis0: block tiles 32 cols × 8 batches per launch. + // ── 8c. (Phase B) Reduce HEAD per-batch grad scratch → final grad buffers. + // Block tree-reduce across bi; deterministic sum order. Each + // launch is grid=(ceil(n_tail/32), 1, 1), block=(32, 8, 1) — + // 32-wide column tile per block, 8-way B-stride per thread. + // 10 launches, one per GRN-head param tensor. AdamW (below) + // reads the final grad buffers; MUST run after these reducers. + // Phase E.3a: encoder reducers (CfC ×4) moved to + // `dispatch_encoder_backward` per + // feedback_single_source_of_truth_no_duplicates. { let n_batch_i = b_sz as i32; - let n_tail_i = HIDDEN_DIM as i32; - let cfg_red = LaunchConfig { - grid_dim: (((HIDDEN_DIM + 31) / 32) as u32, 1, 1), - block_dim: (32, 8, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); - launch - .arg(&self.attn_grad_q_scratch_d) - .arg(&n_batch_i) - .arg(&n_tail_i) - .arg(&mut self.grad_attn_q_d); - unsafe { launch.launch(cfg_red).context("reduce attn_grad_q")?; } - } - - // ── 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] - // Produces: - // grad_x = self.grad_ln_in_d [B, K, H] - // grad_gain/row = self.grad_ln_gain_per_row_d [B*K, H] - // grad_bias/row = self.grad_ln_bias_per_row_d [B*K, H] - // Then two reducer launches collapse the per-row - // scratches into the final [H] param grad buffers - // (no atomicAdd per feedback_no_atomicadd.md). - { - 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_l2_fwd_scratch.h_enriched_seq.cuda_data()) - .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) - .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 (LN_b)")?; } - } - // 7c.i — reduce per-row grad_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_gain_per_row_d) - .arg(&n_rows_ln) - .arg(&mut self.grad_ln_gain_d); - unsafe { launch.launch(cfg_red).context("layer_norm_reduce gain")?; } - } - // 7c.ii — reduce per-row grad_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_bias_per_row_d) - .arg(&n_rows_ln) - .arg(&mut self.grad_ln_bias_d); - unsafe { launch.launch(cfg_red).context("layer_norm_reduce bias")?; } - } - - // ── 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.trunk - .mamba2_l2_mut() - .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.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) - .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.trunk - .mamba2_l1_mut() - .backward_from_h_enriched_seq_full_into( - &self.vsn_out_d, - &self.mamba2_fwd_scratch, - &self.grad_ln_a_in_d, - &mut self.mamba2_bwd_scratch, - &mut self.mamba2_grads_buffers, - ) - .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] - // x = window_tensor_d [B, K, FEATURE_DIM] - // gates = vsn_gates_d (saved by fwd) [B*K, FEATURE_DIM] - // Writes: - // grad_W_vsn, grad_b_vsn — Adam consumes. - // vsn_grad_x_d — discarded (snap_features non-trainable). - // Param grads OVERWRITE (kernel uses += but the prior - // memsets clear them at step start, so a single VSN bwd - // per step writes the full accumulation cleanly). - // Phase B commit 3: VSN per-row grad scratch + reducer. - self.stream.memset_zeros(&mut self.vsn_grad_w_scratch_d) - .map_err(|e| anyhow::anyhow!("zero vsn_grad_w_scratch: {e}"))?; - self.stream.memset_zeros(&mut self.vsn_grad_b_scratch_d) - .map_err(|e| anyhow::anyhow!("zero vsn_grad_b_scratch: {e}"))?; - { - let n_rows_vsn: i32 = (b_sz * k_seq) as i32; - let cfg_vsn_bwd = LaunchConfig { - grid_dim: (n_rows_vsn as u32, 1, 1), - block_dim: (64, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.vsn_bwd_fn); - launch - .arg(&self.trunk.vsn_w_d) - .arg(self.window_tensor_d.cuda_data()) - .arg(&self.vsn_gates_d) - .arg(self.mamba2_grads_buffers.d_x_from_in.cuda_data()) - .arg(&n_rows_vsn) - .arg(&mut self.vsn_grad_w_scratch_d) - .arg(&mut self.vsn_grad_b_scratch_d) - .arg(&mut self.vsn_grad_x_d); - unsafe { launch.launch(cfg_vsn_bwd).context("variable_selection_bwd")?; } - } - // VSN reducer: collapse n_rows (= B * K) → final grad buffers. - // reduce_axis0 tiles 32 columns per block, 8-way batch stride. - { - let n_rows_i = (b_sz * k_seq) as i32; - let n_tail_w_usz = FEATURE_DIM * FEATURE_DIM; - let cfg_red_w = LaunchConfig { - grid_dim: (((n_tail_w_usz + 31) / 32) as u32, 1, 1), - block_dim: (32, 8, 1), - shared_mem_bytes: 0, - }; - let n_tail_w = n_tail_w_usz as i32; - unsafe { - let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); - launch - .arg(&self.vsn_grad_w_scratch_d) - .arg(&n_rows_i) - .arg(&n_tail_w) - .arg(&mut self.grad_vsn_w_d); - launch.launch(cfg_red_w).context("reduce vsn_grad_w")?; - } - let cfg_red_b = LaunchConfig { - grid_dim: (((FEATURE_DIM + 31) / 32) as u32, 1, 1), - block_dim: (32, 8, 1), - shared_mem_bytes: 0, - }; - let n_tail_b = FEATURE_DIM as i32; - unsafe { - let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); - launch - .arg(&self.vsn_grad_b_scratch_d) - .arg(&n_rows_i) - .arg(&n_tail_b) - .arg(&mut self.grad_vsn_b_d); - launch.launch(cfg_red_b).context("reduce vsn_grad_b")?; - } - } - - // ── 8c. (Phase B) Reduce CfC per-batch grad scratch → final grad buffers. - // Block tree-reduce across bi; deterministic sum order. - // Each launch is grid=(ceil(n_tail/32), 1, 1), block=(32, 8, 1) - // — 32-wide column tile per block, 8-way B-stride per thread. - // 4 launches, one per cfc param tensor. AdamW (below) reads - // the final grad buffers; MUST run after these reducers. - { - let n_batch_i = b_sz as i32; - let reduce_at = |n_tail: usize, - scratch: &CudaSlice, - out: &mut CudaSlice, - label: &'static str| + let reduce_heads = |n_tail: usize, + scratch: &CudaSlice, + out: &mut CudaSlice, + label: &'static str| -> Result<()> { let cfg = LaunchConfig { grid_dim: (((n_tail + 31) / 32) as u32, 1, 1), @@ -5271,74 +5570,39 @@ impl PerceptionTrainer { unsafe { launch.launch(cfg).context(label)?; } Ok(()) }; - reduce_at(HIDDEN_DIM * HIDDEN_DIM, &self.cfc_grad_w_in_scratch_d, - &mut self.grad_w_in_d, "reduce cfc_grad_w_in")?; - reduce_at(HIDDEN_DIM * HIDDEN_DIM, &self.cfc_grad_w_rec_scratch_d, - &mut self.grad_w_rec_d, "reduce cfc_grad_w_rec")?; - reduce_at(HIDDEN_DIM, &self.cfc_grad_b_scratch_d, - &mut self.grad_b_d, "reduce cfc_grad_b")?; - reduce_at(HIDDEN_DIM, &self.cfc_grad_tau_scratch_d, - &mut self.grad_tau_d, "reduce cfc_grad_tau")?; // GRN: 10 reducer launches (Phase B commit 2). let nh = N_HORIZONS; let mid = HEAD_MID_DIM; let h = HIDDEN_DIM; - reduce_at(nh * mid * h, &self.grn_grad_w1_scratch_d, + reduce_heads(nh * mid * h, &self.grn_grad_w1_scratch_d, &mut self.grad_heads_w1_d, "reduce grn_grad_w1")?; - reduce_at(nh * mid, &self.grn_grad_b1_scratch_d, + reduce_heads(nh * mid, &self.grn_grad_b1_scratch_d, &mut self.grad_heads_b1_d, "reduce grn_grad_b1")?; - reduce_at(nh * mid * mid, &self.grn_grad_w2_scratch_d, + reduce_heads(nh * mid * mid, &self.grn_grad_w2_scratch_d, &mut self.grad_heads_w2_d, "reduce grn_grad_w2")?; - reduce_at(nh * mid, &self.grn_grad_b2_scratch_d, + reduce_heads(nh * mid, &self.grn_grad_b2_scratch_d, &mut self.grad_heads_b2_d, "reduce grn_grad_b2")?; - reduce_at(nh * mid, &self.grn_grad_w_gate_scratch_d, + reduce_heads(nh * mid, &self.grn_grad_w_gate_scratch_d, &mut self.grad_heads_w_gate_d, "reduce grn_grad_w_gate")?; - reduce_at(nh, &self.grn_grad_b_gate_scratch_d, + reduce_heads(nh, &self.grn_grad_b_gate_scratch_d, &mut self.grad_heads_b_gate_d, "reduce grn_grad_b_gate")?; - reduce_at(nh * mid, &self.grn_grad_w_main_scratch_d, + reduce_heads(nh * mid, &self.grn_grad_w_main_scratch_d, &mut self.grad_heads_w_main_d, "reduce grn_grad_w_main")?; - reduce_at(nh, &self.grn_grad_b_main_scratch_d, + reduce_heads(nh, &self.grn_grad_b_main_scratch_d, &mut self.grad_heads_b_main_d, "reduce grn_grad_b_main")?; - reduce_at(nh * h, &self.grn_grad_w_skip_scratch_d, + reduce_heads(nh * h, &self.grn_grad_w_skip_scratch_d, &mut self.grad_heads_w_skip_d, "reduce grn_grad_w_skip")?; - reduce_at(nh, &self.grn_grad_b_skip_scratch_d, + reduce_heads(nh, &self.grn_grad_b_skip_scratch_d, &mut self.grad_heads_b_skip_d, "reduce grn_grad_b_skip")?; } - // ── 9. Apply AdamW updates on all 17 param groups: CfC×4 + - // GRN heads×10 + LN×2 + Mamba2 grouped. - self.opt_w_in.step(&mut self.trunk.w_in_d, &self.grad_w_in_d)?; - self.opt_w_rec.step(&mut self.trunk.w_rec_d, &self.grad_w_rec_d)?; - self.opt_b.step(&mut self.trunk.b_d, &self.grad_b_d)?; - self.opt_tau.step(&mut self.trunk.tau_all_d, &self.grad_tau_d)?; - // ── Controller B: post-Adam τ projection (spec §3.2) ── - // Per `pearl_adam_normalizes_loss_weights`, the only way to actually - // bound `tau_all_d` is a hard projection AFTER Adam updates the - // parameter. Loss-weight modulation would no-op because Adam's - // m/sqrt(v) cancels constant multipliers. Phase 2 only — Phase 1's - // captured graph never enters this branch since the transition - // invalidates the graph (`pearl_no_host_branches_in_captured_graph`). - if self.phase == TrainingPhase::Phase2Routed { - if let Some(metadata) = self.bucket_routing_metadata.as_ref() { - let cfg_tau_clamp = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (HIDDEN_DIM as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.tau_clamp_fn); - launch - .arg(&mut self.trunk.tau_all_d) - .arg(&metadata.bucket_id_per_channel_d) - .arg(&metadata.bucket_tau_iqr_lo_d) - .arg(&metadata.bucket_tau_iqr_hi_d); - unsafe { - launch - .launch(cfg_tau_clamp) - .context("tau_clamp_kernel launch (Controller B)")?; - } - } - } + // ── 9. Apply HEAD AdamW updates (GRN heads ×10 + Controllers C + + // ALPHA fix). Encoder Adams (CfC ×4 + Controller B + LN ×2 + + // VSN + attn_q + Mamba2 L1/L2) ran inside + // `dispatch_encoder_backward` above. Per + // feedback_single_source_of_truth_no_duplicates the encoder + // Adam steps belong with the encoder backward chain. self.opt_heads_w1.step(&mut self.trunk.heads_w1_d, &self.grad_heads_w1_d)?; self.opt_heads_b1.step(&mut self.trunk.heads_b1_d, &self.grad_heads_b1_d)?; self.opt_heads_w2.step(&mut self.trunk.heads_w2_d, &self.grad_heads_w2_d)?; @@ -5456,13 +5720,6 @@ impl PerceptionTrainer { } } self.opt_heads_b_skip.step(&mut self.trunk.heads_b_skip_d, &self.grad_heads_b_skip_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)?; - self.opt_vsn_b.step(&mut self.trunk.vsn_b_d, &self.grad_vsn_b_d)?; - self.opt_attn_q.step(&mut self.trunk.attn_q_d, &self.grad_attn_q_d)?; // Kendall σ is ISV-driven (closed form from loss_ema in // horizon_ema_and_lambda) — no Adam step. grad_log_sigma_h_d // is still written by the BCE kernel but intentionally ignored. @@ -5471,14 +5728,12 @@ impl PerceptionTrainer { // optimizer groups (horizon_tokens, Q_inv, w_fuse + b_fuse, // moe_gate, experts, log_sigma) will land in commit V10. - // 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(self.trunk.mamba2_l1_mut(), &self.mamba2_grads_buffers) - .context("mamba2 (l1) AdamW step_from_buffers_gpu_clip")?; - self.mamba2_l2_adamw - .step_from_buffers_gpu_clip(self.trunk.mamba2_l2_mut(), &self.mamba2_l2_grads_buffers) - .context("mamba2 (l2) AdamW step_from_buffers_gpu_clip")?; + // NOTE: encoder Adams (LN_b gain/bias, LN_a gain/bias, VSN w/b, + // attn_q, mamba2 L1, mamba2 L2) ran inside + // `dispatch_encoder_backward` above per + // feedback_single_source_of_truth_no_duplicates. The helper is + // the single home for encoder param updates regardless of caller + // (supervised step_batched or RL backward_encoder_with_grad_h_t). // ── SDD-3 Layer B5: aux supervision backward + Adam ── //