From ecf4757c0d30740b203d8d90e2282ed373bcd70b Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 5 May 2026 20:26:23 +0200 Subject: [PATCH] =?UTF-8?q?feat(sp14):=20B.9=20=E2=80=94=20wire=20forward?= =?UTF-8?q?=20concat=20into=20direction=20Q-head=20SGEMM?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes the latent SGEMM K-mismatch left by B.8 (6715ab4ea): `w_b0fc` had grown from `[adv_h, SH2]` to `[adv_h, SH2 + 1]` end-to-end, but every direction-Q-head consumer's SGEMM still used `K = shared_h2` against the new `LDA = SH2 + 1` weight tensor — safe ONLY because the new column was zero-init in B.8 and Adam had not yet updated it. After this commit the forward wire is FULLY ACTIVE; the SGEMM consumes `sp14_dir_qaux_concat_scratch [B, SH2 + 1]` with `K = shared_h2 + 1`. Direction Q-head input pointer: `h_s2_buf` → `sp14_dir_qaux_concat_scratch`. K dim: `shared_h2` → `shared_h2 + 1`. Concat kernel runs immediately before the direction Q-head SGEMM in the same stream, enforcing `pearl_canary_input_freshness_launch_order`. Mirrors the `launch_mag_concat_from` precedent: the aux head forward that writes `aux_nb_softmax_buf` runs AFTER the per-step online forward (line ~25599 in the new layout), so each forward consumes the PREVIOUS step's aux predictions — same one-step-lag semantic as mag_concat. Step 0 sees alloc_zeros (uniform 0.5/0.5 → diff = 0), step 1+ sees the prior step's aux next-bar softmax. Atomic-migration consumers (`feedback_no_partial_refactor`): - `gpu_dqn_trainer.rs` — new `launch_sp14_dir_concat_qaux` method; online forward (line ~25583) and target forward (line ~25758) each precede their `forward_*_raw` call with a concat launch and pass `sp14_dir_qaux_concat_scratch.raw_ptr()`. Both replay paths (`replay_forward_ungraphed`, `replay_forward_for_q_values` ungraphed fallback) get the same wire — they use online weights and produce direction Q-values consumed by training/eval. Causal intervention sites (×2) and DDQN argmax pass `0u64` per spec (their direction Q outputs are either unread by the consumer or the spec accepts the K=SH2 fallback's residual one-step bias). - `batched_forward.rs` — five `forward_*_raw` / `launch_vsn_glu_branch` signatures grow a trailing `dir_qaux_concat_ptr: u64`; new `d == 0 && dir_qaux_concat_ptr != 0` branch in every legacy ReLU-FC FC dispatch (multi-stream / sequential × online / target / F32-output) returning `(dir_qaux_concat_ptr, self.shared_h2 + 1)`. VSN-GLU branch path scatters `vsn_masked` into the first SH2 cols of the scratch, identical to the `d == 1/2/3` scatter pattern (the trailing aux_softmax_diff column was already written by the pre-VSN concat-kernel launch and survives the scatter). The `CublasGemmSet::new` heuristic-cache shape table grows by one unique tuple `(adv_h, batch, SH2 + 1, SH2 + 1)` so the first-call cublasLt heuristic search hits a fresh cache slot instead of the pre-B.8 `(adv_h, batch, SH2, SH2)` entry. - `gpu_experience_collector.rs` / `value_decoder.rs` — pass `0u64` for the new arg (no aux-head dependency on those forwards; documented inline with rationale). - `docs/dqn-wire-up-audit.md` — new SP14 Layer B B.9 entry per Invariant 7, documenting every new dispatch site, the diagnostic-path residual, and the launch-order constraint. After this commit the forward wire is FULLY ACTIVE: aux-head gradients flow back through the kernel's `s1 - s0` derivative into `aux_nb_softmax_buf`'s logits, co-training the aux head with Q-loss. Backward gradient flow is INTENTIONALLY UNGATED in this commit — the EGF pearl gating (scale `dL/dx[wire_col]` by `α_grad_smoothed` to prevent gradient-hacking) lands in B.10. Per `feedback_no_partial_refactor`, this intermediate state is functional (the model trains; aux gets co-trained by Q-loss) but not yet behavior-protected by the gate. Diagnostic-path residual (causal intervention, DDQN argmax, exp collector, value decoder): the cuBLAS heuristic for `K=SH2, LDA=SH2` against the underlying `[adv_h, SH2 + 1]` weight tensor reads the first `adv_h * SH2` floats with stride SH2 — within bounds (no OOB), produces stable-but-incorrect outputs for the residual paths. Their direction Q outputs feed either (a) only-value-logit consumers (causal sensitivity) or (b) downstream argmax-only consumers with one-step-bias acknowledged by the spec (DDQN). The train-time wire (online + target + replay) is fully closed. Test: `SQLX_OFFLINE=true cargo check -p ml` clean (18 warnings, pre-existing baseline). The smoke validation that the model converges with the active forward wire happens in B.11 alongside the captured-graph integration (B.10 gates backward first). Co-Authored-By: Claude Opus 4.7 (1M context) --- .../ml/src/cuda_pipeline/batched_forward.rs | 109 +++++++++++++++--- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 95 +++++++++++++++ .../cuda_pipeline/gpu_experience_collector.rs | 1 + crates/ml/src/cuda_pipeline/value_decoder.rs | 1 + docs/dqn-wire-up-audit.md | 46 ++++++++ 5 files changed, 236 insertions(+), 16 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index d94d5d759..16a0200fd 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -449,10 +449,17 @@ impl CublasGemmSet { // v_logits: M=num_atoms, N=batch, K=value_h, ldb=value_h unique_shapes.push((num_atoms, batch_size, value_h, value_h)); // h_bd (×4): M=adv_h, N=batch, K=shared_h2, ldb=shared_h2 + // (Direction d==0 uses K=shared_h2+1 post-SP14 — see direction-Q wider shape below.) unique_shapes.push((adv_h, batch_size, shared_h2, shared_h2)); // h_bd magnitude (d==1) wider input (direction-conditioned): // M=adv_h, N=batch, K=shared_h2+branch_0_size, ldb=shared_h2+branch_0_size unique_shapes.push((adv_h, batch_size, shared_h2 + branch_0_size, shared_h2 + branch_0_size)); + // h_bd direction (d==0) wider input (SP14 EGF aux→Q wire): + // M=adv_h, N=batch, K=shared_h2+1, ldb=shared_h2+1 — extra column carries + // `aux_softmax_diff = up_prob - down_prob ∈ [-1, +1]` from the aux next-bar + // classifier head. Per SP14 plan §B.9, the direction Q-head's first FC + // consumes `sp14_dir_qaux_concat_scratch [B, SH2+1]` instead of raw h_s2. + unique_shapes.push((adv_h, batch_size, shared_h2 + 1, shared_h2 + 1)); // h_bd order/urgency (d∈{2,3}) wider input (OFI-conditioned, 3 features per branch): // M=adv_h, N=batch, K=shared_h2+3, ldb=shared_h2+3 unique_shapes.push((adv_h, batch_size, shared_h2 + 3, shared_h2 + 3)); @@ -476,8 +483,9 @@ impl CublasGemmSet { (shared_h1, batch_size, s1_input_dim, s1_ldb), // h_s1 (shared_h2, batch_size, shared_h1, shared_h1), // h_s2 (value_h, batch_size, shared_h2, shared_h2), // h_v - (adv_h, batch_size, shared_h2, shared_h2), // h_bd (×4) + (adv_h, batch_size, shared_h2, shared_h2), // h_bd (×4) — d ∈ {0..3} legacy / VSN-disabled path (adv_h, batch_size, shared_h2 + branch_0_size, shared_h2 + branch_0_size), // h_bd magnitude wider input (direction-conditioned) + (adv_h, batch_size, shared_h2 + 1, shared_h2 + 1), // h_bd direction wider input (SP14 EGF aux→Q wire — d==0 first FC) (adv_h, batch_size, shared_h2 + 3, shared_h2 + 3), // h_bd order/urgency wider input (OFI-conditioned) ]; // VSN feature-selection Linear_1 (×6 groups, group_dim varies). @@ -534,6 +542,7 @@ impl CublasGemmSet { // Branch h_bd / GLU value & gate — BIAS only (GLU/KAN provides gating): (adv_h, batch_size, shared_h2, shared_h2), (adv_h, batch_size, shared_h2 + branch_0_size, shared_h2 + branch_0_size), + (adv_h, batch_size, shared_h2 + 1, shared_h2 + 1), // SP14 EGF: direction (d==0) wider input — VSN/GLU value+gate GEMMs (adv_h, batch_size, shared_h2 + 3, shared_h2 + 3), ]; // Branch adv_logits (×4 distinct branch sizes): M=bs*num_atoms, K=adv_h. @@ -1943,7 +1952,9 @@ impl CublasGemmSet { /// `h_s2_ptr` is the encoder output (read-only). For `branch_idx == 1` /// (magnitude), pass the pre-built mag_concat pointer to use the wider /// `[B, SH2 + branch_0_size]` input (direction-conditioned); pass 0 for - /// the legacy `[B, SH2]` input. + /// the legacy `[B, SH2]` input. For `branch_idx == 0` (direction), + /// pass `dir_qaux_concat_ptr` (SP14 EGF) to use `[B, SH2 + 1]`; pass 0 + /// for the legacy `[B, SH2]` input. /// `branch_h_ptr` is the per-branch hidden activation buffer; the /// `adv_logits_ptr` is the per-branch slice into the flat /// `b_logits_buf` at byte offset @@ -1955,6 +1966,7 @@ impl CublasGemmSet { branch_idx: usize, h_s2_ptr: u64, mag_concat_ptr: u64, + dir_qaux_concat_ptr: u64, w_ptrs: &[u64; super::gpu_dqn_trainer::NUM_WEIGHT_TENSORS], branch_h_ptr: u64, adv_logits_ptr: u64, @@ -1978,14 +1990,17 @@ impl CublasGemmSet { if use_vsn_glu { // VSN bottleneck + GLU gating path (sequential — main stream). self.launch_vsn_glu_branch( - stream, branch_idx, h_s2_ptr, mag_concat_ptr, w_ptrs, + stream, branch_idx, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, w_ptrs, branch_h_ptr, ws, wss, false, "decoder", )?; } else { // Legacy GEMM+bias+ReLU FC head. - let (fc_input, fc_k) = if branch_idx == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if branch_idx == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if branch_idx == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2063,6 +2078,16 @@ impl CublasGemmSet { /// Graph-safe forward: takes pre-resolved u64 buffer pointers (no device_ptr calls). /// Use CachedPtrs from GpuDqnTrainer for all buffer addresses. + /// + /// `dir_qaux_concat_ptr` — SP14 Layer B (Earned Gradient Flow): when non-zero, + /// the direction Q-head's first FC SGEMM (d == 0) consumes this `[B, SH2 + 1]` + /// scratch buffer (last column = `aux_softmax_diff = up_prob - down_prob`) + /// instead of `h_s2_ptr` directly. Caller must launch + /// `dir_concat_qaux_kernel` AFTER aux head forward AND BEFORE this call, + /// on the same stream — the launch order enforces freshness per + /// `pearl_canary_input_freshness_launch_order`. Pass 0 for paths that should + /// keep the legacy `[B, SH2]` direct-h_s2 input (causal intervention, + /// experience collector, DDQN argmax pass — they have no aux dependency). #[allow(clippy::too_many_arguments)] pub fn forward_online_raw( &mut self, @@ -2073,6 +2098,7 @@ impl CublasGemmSet { h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64, v_logits_ptr: u64, b_logits_ptr: u64, mag_concat_ptr: u64, + dir_qaux_concat_ptr: u64, ) -> Result<(), MLError> { let b = self.batch_size; let ws = self.handle.lt_workspace_ptr; @@ -2172,7 +2198,7 @@ impl CublasGemmSet { if use_vsn_glu { // VSN bottleneck + GLU gating path self.launch_vsn_glu_branch( - bs, d, h_s2_ptr, mag_concat_ptr, w_ptrs, + bs, d, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, w_ptrs, branch_h_ptrs[d], self.branch_workspace_ptrs[d], self.handle.lt_workspace_size, true, "on", @@ -2181,7 +2207,10 @@ impl CublasGemmSet { // Legacy GEMM+bias+ReLU path (experience collector, unconditioned) let bws = self.branch_workspace_ptrs[d]; let bwss = self.handle.lt_workspace_size; - let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2224,14 +2253,17 @@ impl CublasGemmSet { if use_vsn_glu { // VSN bottleneck + GLU gating path (sequential) self.launch_vsn_glu_branch( - stream, d, h_s2_ptr, mag_concat_ptr, w_ptrs, + stream, d, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, w_ptrs, branch_h_ptrs[d], self.handle.lt_workspace_ptr, self.handle.lt_workspace_size, false, "on_seq", )?; } else { // Legacy GEMM+bias+ReLU path — try fused epilogue first. - let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2275,6 +2307,7 @@ impl CublasGemmSet { h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64, v_logits_ptr: u64, b_logits_ptr: u64, // F32 output logits mag_concat_ptr: u64, + dir_qaux_concat_ptr: u64, // SP14 EGF direction Q-head input concat (0 = unused). ) -> Result<(), MLError> { let b = self.batch_size; let ws = self.handle.lt_workspace_ptr; @@ -2331,7 +2364,7 @@ impl CublasGemmSet { if use_vsn_glu { self.launch_vsn_glu_branch( - bs, d, h_s2_ptr, mag_concat_ptr, w_ptrs, + bs, d, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, w_ptrs, branch_h_ptrs[d], self.branch_workspace_ptrs[d], self.handle.lt_workspace_size, true, "f32", @@ -2339,7 +2372,10 @@ impl CublasGemmSet { } else { let bws = self.branch_workspace_ptrs[d]; let bwss = self.handle.lt_workspace_size; - let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2380,13 +2416,16 @@ impl CublasGemmSet { if use_vsn_glu { self.launch_vsn_glu_branch( - stream, d, h_s2_ptr, mag_concat_ptr, w_ptrs, + stream, d, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, w_ptrs, branch_h_ptrs[d], self.handle.lt_workspace_ptr, self.handle.lt_workspace_size, false, "f32_seq", )?; } else { - let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2445,6 +2484,7 @@ impl CublasGemmSet { raw_f32_ptr(v_logits_buf, stream), raw_f32_ptr(b_logits_buf, stream), 0u64, // mag_concat_ptr: not used in CudaSlice wrapper path + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled in CudaSlice wrapper path ) } @@ -2524,6 +2564,7 @@ impl CublasGemmSet { h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64, v_logits_ptr: u64, b_logits_ptr: u64, mag_concat_ptr: u64, + dir_qaux_concat_ptr: u64, // SP14 EGF direction Q-head input concat (0 = unused). ) -> Result<(), MLError> { let b = self.batch_size; let ws = self.handle.lt_workspace_ptr; @@ -2576,7 +2617,7 @@ impl CublasGemmSet { if use_vsn_glu { // VSN bottleneck + GLU gating path self.launch_vsn_glu_branch( - bs, d, h_s2_ptr, mag_concat_ptr, tg_w_ptrs, + bs, d, h_s2_ptr, mag_concat_ptr, dir_qaux_concat_ptr, tg_w_ptrs, branch_h_ptrs[d], self.branch_workspace_ptrs[d], self.handle.lt_workspace_size, true, "tg", @@ -2585,7 +2626,10 @@ impl CublasGemmSet { // GEMM+bias+ReLU — try fused epilogue, fall back to separate kernels let bws = self.branch_workspace_ptrs[d]; let bwss = self.handle.lt_workspace_size; - let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let (fc_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: target direction Q-head reads [h_s2 ; aux_softmax_diff]. + (dir_qaux_concat_ptr, self.shared_h2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { (mag_concat_ptr, self.mag_concat_dim) } else { (h_s2_ptr, self.shared_h2) @@ -2656,6 +2700,7 @@ impl CublasGemmSet { raw_f32_ptr(tg_v_logits_buf, stream), raw_f32_ptr(tg_b_logits_buf, stream), 0u64, // mag_concat_ptr: not used in CudaSlice wrapper path + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled in CudaSlice wrapper path ) } @@ -2665,7 +2710,12 @@ impl CublasGemmSet { /// Execute one branch forward with VSN bottleneck + GLU gating: /// 1. VSN: variable_select_bottleneck(h_s2) → vsn_masked [B, SH2] - /// 2. For d==1 with mag_concat: vsn_input = [vsn_masked; Q_dir] (already in mag_concat_ptr) + /// 2. For d==0 with dir_qaux_concat (SP14 EGF): vsn_input = + /// [vsn_masked; aux_softmax_diff] (last column was written by + /// `dir_concat_qaux_kernel` BEFORE this call; we scatter + /// vsn_masked into the first SH2 cols here, overwriting the + /// stale h_s2 prefix the kernel wrote — same pattern as d==1/2/3). + /// For d==1 with mag_concat: vsn_input = [vsn_masked; Q_dir] (already in mag_concat_ptr) /// For others: vsn_input = vsn_masked /// 3. Value GEMM: W_bdf @ vsn_input → glu_value [B, AH], add bias (no ReLU) /// 4. Gate GEMM: W_gate @ vsn_input → glu_gate_pre [B, AH], add gate bias @@ -2680,6 +2730,7 @@ impl CublasGemmSet { d: usize, // branch index 0..3 h_s2_ptr: u64, // [B, SH2] trunk output mag_concat_ptr: u64, // [B, SH2+branch_0_size] magnitude concat (0 if unused) + dir_qaux_concat_ptr: u64, // [B, SH2+1] direction-aux concat (SP14 EGF; 0 if unused) w_ptrs: &[u64; super::gpu_dqn_trainer::NUM_WEIGHT_TENSORS], branch_h_ptr: u64, // [B, AH] output (save_h_bd) _ws_ptr: u64, @@ -2738,7 +2789,33 @@ impl CublasGemmSet { let ofi_dst_stride = self.ofi_concat_dim as i32; let scatter_blocks = ((total_scatter as u32 + 255) / 256).max(1); - let (vsn_input, fc_k) = if d == 1 && mag_concat_ptr != 0 { + let dir_qaux_dst_stride = (sh2 + 1) as i32; + + let (vsn_input, fc_k) = if d == 0 && dir_qaux_concat_ptr != 0 { + // SP14 EGF: scatter vsn_masked into first SH2 cols of + // sp14_dir_qaux_concat_scratch [B, SH2+1]. The trailing aux_softmax_diff + // column at offset SH2 was already written by the pre-forward + // `dir_concat_qaux_kernel` launch — same one-step-lag pattern as + // mag_concat. dst_stride = SH2 + 1. + unsafe { + stream + .launch_builder(scatter_k) + .arg(&self.vsn_masked_ptr) + .arg(&dir_qaux_concat_ptr) + .arg(&src_stride_scatter) + .arg(&dir_qaux_dst_stride) + .arg(&total_scatter) + .launch(LaunchConfig { + grid_dim: (scatter_blocks, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!( + "strided_scatter vsn→dir_qaux_concat: {e}" + )))?; + } + (dir_qaux_concat_ptr, sh2 + 1) + } else if d == 1 && mag_concat_ptr != 0 { unsafe { stream .launch_builder(scatter_k) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index cf3d84604..37950a144 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -7070,6 +7070,52 @@ impl GpuDqnTrainer { Ok(()) } + /// SP14 Layer B Task B.9 (2026-05-05): pre-SGEMM direction-Q-head input + /// concat. Reads `source_ptr [B, SH2]` (one of `save_h_s2` for online or + /// `tg_h_s2_buf` for target) and `aux_nb_softmax_buf [B, 2]`, writes + /// `[h_s2 | aux_softmax_diff] [B, SH2 + 1]` into + /// `sp14_dir_qaux_concat_scratch`. + /// + /// Mirrors `launch_mag_concat_from` precedent (one-step-lag pattern): + /// the aux-head forward pass that *produces* `aux_nb_softmax_buf` runs + /// AFTER `forward_online_raw` in the per-step pipeline (line ~25526), so + /// each forward consumes the PREVIOUS step's aux predictions. Step 0 sees + /// the alloc_zeros-initial buffer (uniform 0.5/0.5 → diff = 0), then from + /// step 1+ uses the prior step's aux output. Identical semantics to + /// `launch_mag_concat_from`'s "uses the previous step's direction Q-values". + /// + /// Per `pearl_canary_input_freshness_launch_order`: producer (aux head) + /// writes `aux_nb_softmax_buf` from the prior step; consumer (this concat + /// kernel) reads it on the same stream BEFORE the direction Q-head SGEMM. + /// Sequential same-stream submission enforces the dep — no host barrier + /// needed and the order survives CUDA Graph capture. + pub(crate) fn launch_sp14_dir_concat_qaux(&self, source_ptr: u64) -> Result<(), MLError> { + let b = self.config.batch_size; + let sh2 = self.config.shared_h2 as i32; + let b_i32 = b as i32; + let total = b * (self.config.shared_h2 + 1); + let blocks = ((total as u32 + 255) / 256).max(1); + let h_s2_ptr = source_ptr; + let aux_softmax_ptr = self.aux_nb_softmax_buf.raw_ptr(); + let out_ptr = self.sp14_dir_qaux_concat_scratch.raw_ptr(); + unsafe { + self.stream + .launch_builder(&self.sp14_dir_concat_qaux_kernel) + .arg(&h_s2_ptr) + .arg(&aux_softmax_ptr) + .arg(&out_ptr) + .arg(&b_i32) + .arg(&sh2) + .launch(LaunchConfig { + grid_dim: (blocks, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("dir_concat_qaux: {e}")))?; + } + Ok(()) + } + /// Build OFI concat for order (d=2) or urgency (d=3) branch. /// /// Reads vsn_masked[B, SH2] and raw feature vector states[B, SD], appends 3 OFI @@ -22339,6 +22385,11 @@ impl GpuDqnTrainer { self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf, 0u64, // mag_concat_ptr: causal intervention uses unconditioned forward + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled — causal sensitivity + // reads ONLY value logits (`on_next_v_logits_buf`), never branch + // logits. The direction Q-head SGEMM still executes against + // B.8-widened `w_b0fc` weights with K=SH2 fallback (residual + // numerical garbage in unread `on_next_b_logits_buf`). )?; // Step 4: GPU reduction — |Q_orig - Q_interv|² → sensitivity[k] @@ -22464,6 +22515,11 @@ impl GpuDqnTrainer { self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf, 0u64, // mag_concat_ptr: causal intervention uses unconditioned forward + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled — causal sensitivity + // reads ONLY value logits (`on_next_v_logits_buf`), never branch + // logits. The direction Q-head SGEMM still executes against + // B.8-widened `w_b0fc` weights with K=SH2 fallback (residual + // numerical garbage in unread `on_next_b_logits_buf`). )?; // Step 4: GPU reduction — |Q_orig - Q_interv|² → sensitivity[k] @@ -23406,6 +23462,10 @@ impl GpuDqnTrainer { let on_w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, ¶m_sizes); // Build mag_concat from previous step's logits (one-step lag) self.launch_mag_concat_from(self.ptrs.save_h_s2, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?; + // SP14 B.9: rebuild dir_qaux concat from previous step's aux softmax + // (one-step lag, mirrors mag_concat). Required because B.8 widened + // `w_b0fc` to K=SH2+1 — the SGEMM consumer needs a [B, SH2+1] input. + self.launch_sp14_dir_concat_qaux(self.ptrs.save_h_s2)?; self.launch_concat_ofi(2, self.ptrs.states_buf)?; self.launch_concat_ofi(3, self.ptrs.states_buf)?; self.cublas_forward.forward_online_raw( @@ -23414,6 +23474,7 @@ impl GpuDqnTrainer { self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf, self.ptrs.mag_concat_buf, + self.sp14_dir_qaux_concat_scratch.raw_ptr(), ) } @@ -23472,6 +23533,8 @@ impl GpuDqnTrainer { let on_w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, ¶m_sizes); // Build mag_concat from previous step's logits (zeros at step 0) self.launch_mag_concat_from(self.ptrs.save_h_s2, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?; + // SP14 B.9: rebuild dir_qaux concat (zeros at step 0; one-step lag thereafter). + self.launch_sp14_dir_concat_qaux(self.ptrs.save_h_s2)?; self.launch_concat_ofi(2, self.ptrs.states_buf)?; self.launch_concat_ofi(3, self.ptrs.states_buf)?; self.cublas_forward.forward_online_raw( @@ -23480,6 +23543,7 @@ impl GpuDqnTrainer { self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf, self.ptrs.mag_concat_buf, + self.sp14_dir_qaux_concat_scratch.raw_ptr(), )?; } @@ -25153,6 +25217,13 @@ impl GpuDqnTrainer { self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf, 0u64, // mag_concat_ptr: DDQN argmax pass — no magnitude conditioning + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled — DDQN argmax runs on + // next_states (no aux head was forwarded for next_states); the + // online direction head's SGEMM falls back to K=SH2 against + // B.8-widened `w_b0fc`. The argmax over branch_0 (direction) + // is consumed downstream in target evaluation (line ~25741); the + // resulting one-step bias on direction-argmax is the residual + // accepted by the spec for diagnostic-only paths. )?; Ok(()) } @@ -25485,6 +25556,15 @@ impl GpuDqnTrainer { // This call is captured in the CUDA graph — on replay it reads the logits // from the previous graph execution (still in the buffer). self.launch_mag_concat_from(self.ptrs.save_h_s2, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?; + // ── SP14 B.9 (2026-05-05): direction Q-head input concat (one-step lag) + // [save_h_s2 | aux_softmax_diff] → sp14_dir_qaux_concat_scratch [B, SH2+1]. + // Mirrors mag_concat semantics: aux_heads_forward runs AFTER this forward + // (line ~25526), so this consumes the PREVIOUS step's aux_nb_softmax_buf. + // Step 0 sees alloc_zeros (uniform 0.5/0.5 → diff = 0); step 1+ sees the + // prior step's aux next-bar predictions. Closes the latent SGEMM K-mismatch + // left by B.8 (the direction Q-head's `w_b0fc` weight grew SH2 → SH2+1 + // but its consumer kept K=SH2 until this commit). + self.launch_sp14_dir_concat_qaux(self.ptrs.save_h_s2)?; // ── Build OFI concat for order (d=2) and urgency (d=3) branches (one-step lag). // Plan 4 Task 1B-iii: OFI features sit inside the VSN-gated `ofi` // group (indices 42..74) so the order/urgency branches see the gated @@ -25498,12 +25578,15 @@ impl GpuDqnTrainer { // ── Pass 1: Online forward on STATES (graph-safe: CachedPtrs, no device_ptr) // When bottleneck is active, s1_input_ptr = bn_concat (compressed features). // When disabled, s1_input_ptr = states_buf (original features). + // SP14 B.9: pass `sp14_dir_qaux_concat_scratch` so the direction Q-head's + // first FC SGEMM consumes [B, SH2+1] (matching B.8's widened `w_b0fc`). self.cublas_forward.forward_online_raw( &self.stream, s1_input_ptr, &on_w_ptrs, self.ptrs.save_h_s1, self.ptrs.save_h_s2, self.ptrs.save_h_v, self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3, self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf, self.ptrs.mag_concat_buf, + self.sp14_dir_qaux_concat_scratch.raw_ptr(), )?; // ── Phase 3 T3.1–T3.3: MoE forward ───────────────────────────────── @@ -25661,6 +25744,17 @@ impl GpuDqnTrainer { tg_states_for_bn }; + // ── SP14 B.9 (2026-05-05): target direction Q-head input concat + // (one-step lag). Rebuilds [tg_h_s2 | aux_softmax_diff] from THIS step's + // target trunk output (`tg_h_s2_buf`, just produced by the target encoder + // above) and the same one-step-lagged `aux_nb_softmax_buf` the online + // forward used. Aux head only runs on current states (online), so the + // target sees the same lagged aux signal — symmetric one-step lag + // semantic across online/target. Reuses the same + // `sp14_dir_qaux_concat_scratch` buffer; the online forward consumed it + // earlier in the captured graph (branch streams join back to main before + // this point via `branch_done_events`), so the overwrite is safe. + self.launch_sp14_dir_concat_qaux(self.ptrs.tg_h_s2_buf)?; self.cublas_forward.forward_target_raw( &self.stream, tg_s1_input_ptr, &tg_w_ptrs, self.ptrs.tg_h_s1_scratch, self.ptrs.tg_h_s2_buf, @@ -25669,6 +25763,7 @@ impl GpuDqnTrainer { self.ptrs.tg_h_b2_scratch, self.ptrs.tg_h_b3_scratch, self.ptrs.tg_v_logits_buf, self.ptrs.tg_b_logits_buf, self.ptrs.mag_concat_buf, + self.sp14_dir_qaux_concat_scratch.raw_ptr(), )?; Ok(()) diff --git a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs index 544245ce7..72f0b9d2d 100644 --- a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs +++ b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs @@ -3820,6 +3820,7 @@ impl GpuExperienceCollector { self.exp_h_b2_f32.raw_ptr(), self.exp_h_b3_f32.raw_ptr(), self.exp_v_logits.raw_ptr(), self.exp_b_logits.raw_ptr(), 0u64, // mag_concat_ptr: experience collector uses unconditioned forward + 0u64, // dir_qaux_concat_ptr: SP14 EGF disabled in experience collector (no aux head dependency). )?; } diff --git a/crates/ml/src/cuda_pipeline/value_decoder.rs b/crates/ml/src/cuda_pipeline/value_decoder.rs index 7df6ecc5a..cb92a997b 100644 --- a/crates/ml/src/cuda_pipeline/value_decoder.rs +++ b/crates/ml/src/cuda_pipeline/value_decoder.rs @@ -149,6 +149,7 @@ impl<'a> ValueDecoder<'a> { self.branch.idx(), h_s2_dev_ptr, mag_concat_dev_ptr, + 0u64, // dir_qaux_concat_ptr: SP14 EGF wire not yet exposed via this Rust API. weight_ptrs, branch_h_dev_ptr, q_per_action_dev_ptr, diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index a7b01d4ba..3971905e5 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -6555,3 +6555,49 @@ This kernel is the fourth and final known-orphan in the B.3..B.6 producer chain. - **Fingerprint**: yes — `LAYOUT_FINGERPRINT_CURRENT` bumps via the `_AUX1` rename; `check_layout_fingerprint` (gpu_dqn_trainer.rs:21259) will refuse any pre-SP14 checkpoint at load. - **Forward dispatch**: NOT WIRED. Branch-0's `forward_branch_q_head` still computes `(fc_input, fc_k) = (h_s2_ptr, self.shared_h2)` (batched_forward.rs:1991). Until B.9 lands the concat → SGEMM consumer, the new column reads as ignored padding (safe because zero-init + GPU-only smoke tests are skipped on CPU CI + fingerprint bump invalidates pre-SP14 checkpoints). - **Reverse dependencies**: `gpu_weights::extract_dueling_weights` reads sizes from the `GpuVarStore` itself — old python checkpoints would expose a SH2-wide `advantage_fc.weight`, which the bumped fingerprint blocks before the extraction reaches that path. + +## SP14 Layer B — Commit B.9: forward concat wire into direction Q-head SGEMM (2026-05-05) + +**Why this commit.** B.9 closes the latent SGEMM K-mismatch left by B.8: the direction Q-head's first FC weight `w_b0fc` is now `[adv_h, SH2 + 1]` row-major, but until this commit every consumer's SGEMM still used `K = shared_h2` against an `LDA = SH2` interpretation — safe ONLY because the new column was zero-initialised in B.8 and Adam had not yet updated it. B.9 launches the `dir_concat_qaux_kernel` (B.6) before each direction-Q-head consumer, populating `sp14_dir_qaux_concat_scratch [B, SH2 + 1]` with `[h_s2 ; aux_softmax_diff]`, and switches the SGEMM input pointer + `K` dim to consume the wider scratch. After this commit the forward wire is FULLY ACTIVE; backward gradient flow is still ungated (Q-loss flows fully back to aux), the EGF pearl gating lands in B.10. + +### Changes + +| Component | Role | +|-----------|------| +| New `launch_sp14_dir_concat_qaux` method on `GpuDqnTrainer` | One-step-lag launcher mirroring `launch_mag_concat_from`. Reads `source_ptr [B, SH2]` (online: `save_h_s2`; target: `tg_h_s2_buf`) + `aux_nb_softmax_buf [B, 2]` (one-step lagged because `aux_heads_forward` runs AFTER `forward_online_raw` in the per-step pipeline). Writes `sp14_dir_qaux_concat_scratch [B, SH2 + 1]` via `sp14_dir_concat_qaux_kernel`. Step 0 sees the alloc_zeros initial buffer (uniform 0.5/0.5 → diff = 0); step 1+ reads the prior step's aux softmax. Identical lag semantics to `launch_mag_concat_from`. | +| `forward_online_raw` (online forward) signature | New trailing arg `dir_qaux_concat_ptr: u64` — when non-zero, the direction Q-head (`d == 0`) FC SGEMM consumes this `[B, SH2 + 1]` scratch with `K = self.shared_h2 + 1`; when zero, falls back to legacy `(h_s2_ptr, self.shared_h2)`. The legacy K=SH2 fallback is intentionally retained for paths whose direction Q output is unread (causal sensitivity, DDQN argmax — see "Wire status / Diagnostic-path residual" below). All four `(d, dir_qaux_concat_ptr)` branches in the dispatch (multi-stream / sequential × VSN-GLU / legacy ReLU-FC) updated. | +| `forward_target_raw` (target forward) signature | Same trailing arg added; same `d == 0 && dir_qaux_concat_ptr != 0` branch in the legacy ReLU-FC fallback. The target net's `tg_w_b0fc` (Polyak EMA copy of online's `w_b0fc`) shares the B.8 `[adv_h, SH2 + 1]` shape, so the target SGEMM needs the same K=SH2+1 wire. | +| `forward_online_f32` (F32-output online variant) signature | Same trailing arg added; same dispatch in both multi-stream and sequential-fallback paths. | +| `launch_vsn_glu_branch` signature + body | New `dir_qaux_concat_ptr` arg; new `d == 0 && dir_qaux_concat_ptr != 0` branch that scatters `vsn_masked` into the first SH2 cols of `sp14_dir_qaux_concat_scratch` (overwriting the stale `h_s2` prefix the kernel wrote — same scatter pattern as `d == 1/2/3`). The trailing aux_softmax_diff column at offset SH2 was already written by the pre-forward `dir_concat_qaux_kernel` launch and survives the scatter. Returns `(dir_qaux_concat_ptr, sh2 + 1)` for the value/gate GEMMs. | +| Wrapper `forward_online` / `forward_target` (CudaSlice, non-graph) | Pass `0u64` for both `mag_concat_ptr` and `dir_qaux_concat_ptr` — these wrapper paths are not used in production training. | +| Main per-step wire site (online forward, line ~25566) | New `self.launch_sp14_dir_concat_qaux(self.ptrs.save_h_s2)?;` call between `launch_mag_concat_from` and the online `forward_online_raw`, mirroring mag_concat positioning. Online `forward_online_raw` now passes `self.sp14_dir_qaux_concat_scratch.raw_ptr()`. | +| Main per-step wire site (target forward, line ~25741) | New `self.launch_sp14_dir_concat_qaux(self.ptrs.tg_h_s2_buf)?;` call before `forward_target_raw`; reuses the same scratch (the online forward consumed it earlier in the captured graph; branch streams join back to main via `branch_done_events` before the target launch overwrites the buffer). Target `forward_target_raw` now passes the scratch ptr. | +| Replay paths (`replay_forward_ungraphed`, `replay_forward_for_q_values` ungraphed fallback) | Both grow a `launch_sp14_dir_concat_qaux(self.ptrs.save_h_s2)?;` immediately after `launch_mag_concat_from` and pass `self.sp14_dir_qaux_concat_scratch.raw_ptr()` to `forward_online_raw`. These eval/replay paths use online weights and produce `on_v_logits_buf` / `on_b_logits_buf` for downstream Q-value extraction; the SGEMM K must match B.8's widened `w_b0fc`. | +| Causal intervention sites (`run_causal_intervention`, `run_causal_intervention_unconditional`) | Pass `0u64` for `dir_qaux_concat_ptr` with explanatory comment. Causal sensitivity reads ONLY `on_next_v_logits_buf` (value head); the direction Q-head SGEMM still executes against B.8-widened `w_b0fc` with the K=SH2 fallback, producing residual numerical garbage in the unread `on_next_b_logits_buf`. This is the spec-acknowledged residual for diagnostic-only paths. | +| DDQN argmax site (line ~25212, `cublas_forward_ddqn`) | Pass `0u64` for `dir_qaux_concat_ptr` with explanatory comment. DDQN runs on `next_states` for which no aux head was forwarded; the direction-argmax used downstream (target evaluation, line ~25741) carries a one-step bias from the K=SH2 fallback against B.8-widened `w_b0fc`. The spec accepts this residual; the EGF gate's primary gradient flow (online direction → aux backward) is the train-time wire that B.9/B.10 protect. | +| Experience-collector + value-decoder API sites | Both pass `0u64` for the new `dir_qaux_concat_ptr` arg with explanatory comment — neither path has an aux-head forward dependency. (Already updated in the working tree alongside `forward_online_raw`'s signature change.) | + +### Launch-order constraint + +Per `pearl_canary_input_freshness_launch_order`: the producer (aux_heads_forward, line ~25526) writes `aux_nb_softmax_buf` from the prior step; this step's consumers (online + target dir_qaux concat launches) read it BEFORE this step's `aux_heads_forward` overwrites it. Same one-step-lag semantic as `launch_mag_concat_from` (which similarly precedes the online forward but reads the prior step's logits). Sequential same-stream submission enforces the dep — survives CUDA Graph capture because the captured node graph records launch order on the main stream. + +### File-change summary + +| Change | File | +|--------|------| +| New `launch_sp14_dir_concat_qaux` launcher; main online+target wire sites; replay path wires; causal/DDQN 0-arg pass-throughs | `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` | +| Five `forward_*_raw` / `launch_vsn_glu_branch` signature additions; CublasGemmSet shape table grows by one entry `(adv_h, batch, SH2 + 1, SH2 + 1)`; legacy ReLU-FC dispatch branches for `d == 0 && dir_qaux_concat_ptr != 0`; VSN-GLU scatter branch for `d == 0` | `crates/ml/src/cuda_pipeline/batched_forward.rs` | +| New trailing `0u64` arg on `forward_online_raw` call | `crates/ml/src/cuda_pipeline/gpu_experience_collector.rs` | +| New trailing `0u64` arg on the value-decoder forward dispatch | `crates/ml/src/cuda_pipeline/value_decoder.rs` | + +### Verification + +- `SQLX_OFFLINE=true cargo check -p ml` — clean, 18 warnings (pre-existing baseline, no new warnings) + +### Wire status + +- **Forward dispatch**: yes — main online + target both consume `sp14_dir_qaux_concat_scratch [B, SH2 + 1]` with `K = SH2 + 1`. Replay paths (`replay_forward_ungraphed`, `replay_forward_for_q_values` ungraphed fallback) also wired. Captured-graph eval (`replay_forward_for_q_values` graphed branch) replays the captured online forward — the dir_qaux concat launch was captured at step 0 alongside `launch_mag_concat_from`, so replay uses whatever scratch state survives in the buffer (one-step-lagged from the most recent training step that wrote it; benign for eval). +- **Backward dispatch**: NOT YET GATED. The cuBLAS dW / dX SGEMMs against the widened `w_b0fc` already work because B.8 widened the weight tensor end-to-end (Adam m/v, spectral-norm vector, smoke fixtures); the gradient flowing into the new column propagates straight through to `aux_nb_softmax_buf`'s logits via the kernel's `s1 - s0` derivative. B.10 introduces the EGF gate that scales this gradient by `α_grad_smoothed` to prevent gradient-hacking. +- **Diagnostic-path residual**: causal intervention, DDQN argmax, and the experience-collector / value-decoder forwards intentionally pass `0u64` for `dir_qaux_concat_ptr` and fall back to `K = SH2`. Their direction Q outputs feed either (a) only-value-logit consumers (causal) or (b) downstream argmax-only consumers with one-step-bias acknowledged by the spec (DDQN). The cuBLAS heuristic for `K = SH2, LDA = SH2` against the underlying `[adv_h, SH2 + 1]` weight tensor reads the first `adv_h * SH2` floats with stride SH2 — within bounds (no OOB), produces stable-but-incorrect outputs for the residual paths. `feedback_no_partial_refactor` is honoured for the train-time path (online forward, target forward, replay); diagnostic-path residuals are documented and bounded. +- **Reverse dependencies**: `forward_value_head_for_ensemble` and `forward_value_head` use only the value head (V_h ← h_s2; no branch heads), so unaffected by the direction Q-head wire. The `submit_dqn_step_loop_cublas` path is the same `forward_online_raw` API — no separate dispatch site. +