diff --git a/crates/ml/build.rs b/crates/ml/build.rs index 90b3c2a44..408e5ec74 100644 --- a/crates/ml/build.rs +++ b/crates/ml/build.rs @@ -206,6 +206,18 @@ fn main() { // Phase C.4. See plan §C.3 in // `2026-05-07-sp14-layer-c-separate-aux-trunk.md`. "aux_trunk_forward_kernel.cu", + // SP14 Layer C Phase C.4 (2026-05-08): aux trunk backward kernel — + // three kernels in one cubin (`aux_trunk_bwd_dh_pre`, + // `aux_trunk_bwd_dW_reduce`, `aux_trunk_bwd_db_reduce`). + // Propagates `dh_s2_aux` through w3/w2/w1 to accumulate + // dW{1,2,3} + db{1,2,3} via block-tree-reduce (no atomicAdd per + // feedback_no_atomicadd). CRITICAL: kernel does NOT compute + // `dx_in` — the encoder boundary is the stop-gradient per + // locked design. Module is additive — wire-up into collector + // backward chain + Adam updates land in Phase C.5 (atomic). + // See plan §C.4 in + // `2026-05-07-sp14-layer-c-separate-aux-trunk.md`. + "aux_trunk_backward_kernel.cu", // Plan A Phase 0 (Thompson sampling spec 2026-04-26): standalone test // kernel exercised only by `distributional_q_tests.rs`. Implements the // Thompson direction sampling math (inverse-CDF over C51 atoms, diff --git a/crates/ml/src/cuda_pipeline/aux_trunk_backward_kernel.cu b/crates/ml/src/cuda_pipeline/aux_trunk_backward_kernel.cu new file mode 100644 index 000000000..ab1246de7 --- /dev/null +++ b/crates/ml/src/cuda_pipeline/aux_trunk_backward_kernel.cu @@ -0,0 +1,287 @@ +/* + * aux_trunk_backward_kernel.cu — SP14 Layer C Phase C.4 (2026-05-08). + * + * Backward pass for the auxiliary trunk's 3-layer Linear→ELU→Linear→ELU→ + * Linear MLP introduced in Phase C.3 (`aux_trunk_forward_kernel.cu`). + * Receives `dh_s2_aux [B, AUX_HIDDEN_DIM]` from the aux head's backward + * and propagates through w3/w2/w1 to accumulate parameter gradients + * `dW{1,2,3}` and bias gradients `db{1,2,3}`. + * + * Forward chain (recap from C.3): + * + * x_in [B, ENCODER_OUT_DIM] + * → Linear(w1, b1) → h_aux1_pre → ELU → h_aux1 [B, H1] (saved-fwd) + * → Linear(w2, b2) → h_aux2_pre → ELU → h_aux2 [B, H2] (saved-fwd) + * → Linear(w3, b3) → h_s2_aux [B, AUX_HIDDEN_DIM] (no act) + * + * Backward math (in reverse order): + * + * d_logits = dh_s2_aux (output is linear) + * db3[j] = sum_b d_logits[b, j] + * dW3[k, j] = sum_b h_aux2[b, k] * d_logits[b, j] + * dh_aux2_post[b, k] = sum_j d_logits[b, j] * w3[k, j] + * dh_aux2_pre [b, k] = dh_aux2_post[b, k] * elu_bwd_post(h_aux2[b, k]) + * + * db2[j] = sum_b dh_aux2_pre[b, j] + * dW2[k, j] = sum_b h_aux1[b, k] * dh_aux2_pre[b, j] + * dh_aux1_post[b, k] = sum_j dh_aux2_pre[b, j] * w2[k, j] + * dh_aux1_pre [b, k] = dh_aux1_post[b, k] * elu_bwd_post(h_aux1[b, k]) + * + * db1[j] = sum_b dh_aux1_pre[b, j] + * dW1[k, j] = sum_b x_in[b, k] * dh_aux1_pre[b, j] + * + * STOP — NO `dx_in` write. The encoder boundary is the stop-gradient + * per the locked design decision (see plan §C.4 "Critical: NO + * `dx_in_out` parameter"). The aux trunk reshapes its own representation + * from x_in but cannot pull on the encoder. Q-loss remains the sole + * shaping force on the encoder. + * + * Pearls applied: + * - feedback_no_atomicadd.md — every dW/db output element is owned by + * exactly one block; the block performs a shmem-tree reduction over + * the batch dim. No atomicAdd anywhere. + * - feedback_no_stubs.md — full SGEMM-style accumulation, no + * placeholder body, no missing layer. + * - pearl_no_host_branches_in_captured_graph.md — kernel signatures are + * pure-device; the wrapper's `launch()` issues the fixed sequence of + * kernel launches with no host-side branching. + * + * Topology / layout (mirrors C.3 forward; passed at launch as runtime args): + * - ENCODER_OUT_DIM = config.shared_h1 (256 in production) + * - H1 = AUX_TRUNK_H1 (256) + * - H2 = AUX_TRUNK_H2 (128) + * - AUX_HIDDEN_DIM = config.shared_h2 (256) + * + * w1 [ENCODER_OUT_DIM, H1] row-major → w1[k, j] = w1_ptr[k * H1 + j] + * w2 [H1, H2] row-major → w2[k, j] = w2_ptr[k * H2 + j] + * w3 [H2, AUX_HIDDEN_DIM] row-major → w3[k, j] = w3_ptr[k * AUX_HIDDEN_DIM + j] + * + * Saved-fwd buffers (POST-ELU values from C.3 forward): + * h_aux1 [B, H1] — Layer-1 post-ELU + * h_aux2 [B, H2] — Layer-2 post-ELU + * + * ELU' from POST-activation `y = ELU(x)` per `aux_elu_bwd_from_post` in + * `aux_heads_kernel.cu`: + * if y > 0: x > 0 ⇒ ELU'(x) = 1 + * else: x ≤ 0 ⇒ ELU(x) = exp(x) - 1, so ELU'(x) = exp(x) = 1 + y + * (the y > 0 branch covers both x = 0 and x > 0 — at x = 0, y = 0 so + * the ≤ branch fires and returns 1.0 anyway: ELU is C¹ at the origin.) + * + * Kernel sequence (orchestrated by `AuxTrunkBackwardOps::launch`): + * + * 1. aux_trunk_bwd_dh_pre — [B] blocks, AUX_TRUNK_BLOCK threads. + * in: dh_s2_aux [B, AUX_HIDDEN_DIM], w3, w2, h_aux1, h_aux2 + * out: dh_aux2_pre [B, H2], dh_aux1_pre [B, H1] (scratch) + * + * 2. aux_trunk_bwd_dW_reduce — generic dW = sum_b A[b, k] * B[b, j]. + * Grid: (Krows * Jcols, 1, 1). Block-tree-reduce over batch. + * Launched 3× (dW3, dW2, dW1) with different (A, B) input pairs. + * + * 3. aux_trunk_bwd_db_reduce — generic db = sum_b B[b, j]. + * Grid: (J, 1, 1). Block-tree-reduce over batch. + * Launched 3× (db3, db2, db1). + * + * No layer4 / dx_in kernel exists by design — that's the stop-gradient. + */ + +#include +#include + +#ifndef AUX_TRUNK_BWD_BLOCK +#define AUX_TRUNK_BWD_BLOCK 256 +#endif + +/* ELU backward from POST-activation y = ELU(x). Mirrors + * `aux_elu_bwd_from_post` in aux_heads_kernel.cu — kept inline because the + * common header pre-pended by build.rs only ships generic device fns; the + * aux-specific ELU helpers live in their own kernel files. */ +__device__ __forceinline__ float aux_trunk_elu_bwd_from_post(float y) { + return (y > 0.0f) ? 1.0f : (1.0f + y); +} + +/* ===================================================================== + * aux_trunk_bwd_dh_pre — pre-pass for the parameter-gradient reduces. + * + * Per sample b (one block per sample, AUX_TRUNK_BWD_BLOCK threads): + * 1. Compute dh_aux2_pre[b, :] = (d_logits[b, :] @ w3.T) ⊙ ELU'(h_aux2[b, :]). + * Stored to global `dh_aux2_pre_out [B, H2]` for the dW2/db2/dW1 + * reduces. Cached in shmem `sh_dh2_pre[H2]` for the next layer's + * matmul. + * 2. Compute dh_aux1_pre[b, :] = (sh_dh2_pre @ w2.T) ⊙ ELU'(h_aux1[b, :]). + * Stored to global `dh_aux1_pre_out [B, H1]` for dW1/db1. + * + * Note: each block computes both layers serially within itself (no inter- + * block communication required because both layers consume the same + * sample's data). The shmem `sh_dh2_pre[H2]` lives across layers so + * Layer-2 matmul reads from shmem rather than a second global-memory + * round-trip. + * + * Block: AUX_TRUNK_BWD_BLOCK = 256 threads. Grid: (B, 1, 1). + * Shared memory: H2 floats (Layer-2 post-multiply cache). + * + * No atomicAdd, no inter-block contention — each (b, k) output cell is + * owned by exactly one block, written by exactly one thread. + * ===================================================================== */ +extern "C" __global__ void aux_trunk_bwd_dh_pre( + const float* __restrict__ d_logits, /* [B, AUX_HIDDEN_DIM] */ + const float* __restrict__ w3, /* [H2, AUX_HIDDEN_DIM] row-major */ + const float* __restrict__ w2, /* [H1, H2] row-major */ + const float* __restrict__ h_aux1, /* [B, H1] post-ELU */ + const float* __restrict__ h_aux2, /* [B, H2] post-ELU */ + float* __restrict__ dh_aux2_pre_out, /* [B, H2] OUT */ + float* __restrict__ dh_aux1_pre_out, /* [B, H1] OUT */ + int B, + int H1, + int H2, + int AUX_HIDDEN_DIM +) { + const int b = blockIdx.x; + if (b >= B) return; + const int tid = threadIdx.x; + + /* Shared memory: H2-vector cache for Layer-2 post-multiply (so the + * Layer-1 dh_aux1 matmul reads from shmem). */ + extern __shared__ float smem[]; + float* sh_dh2_pre = smem; /* [H2] */ + + /* ─────────────── Layer 3 backward → dh_aux2_pre ─────────────── */ + /* Each thread owns a slice of H2 via stride loop. Serial dot over + * AUX_HIDDEN_DIM per lane. */ + { + const float* dlogits_row = d_logits + (size_t)b * AUX_HIDDEN_DIM; + for (int k = tid; k < H2; k += blockDim.x) { + float acc = 0.0f; + const float* w3_row = w3 + (size_t)k * AUX_HIDDEN_DIM; + for (int j = 0; j < AUX_HIDDEN_DIM; ++j) { + acc += dlogits_row[j] * w3_row[j]; + } + const float h2_post = h_aux2[(size_t)b * H2 + k]; + const float dh2_pre = acc * aux_trunk_elu_bwd_from_post(h2_post); + sh_dh2_pre[k] = dh2_pre; + dh_aux2_pre_out[(size_t)b * H2 + k] = dh2_pre; + } + } + __syncthreads(); + + /* ─────────────── Layer 2 backward → dh_aux1_pre ─────────────── */ + /* dh_aux1_post[b, k] = sum_j sh_dh2_pre[j] * w2[k, j] + * dh_aux1_pre [b, k] = dh_aux1_post[b, k] * ELU'(h_aux1[b, k]) */ + for (int k = tid; k < H1; k += blockDim.x) { + float acc = 0.0f; + const float* w2_row = w2 + (size_t)k * H2; + for (int j = 0; j < H2; ++j) { + acc += sh_dh2_pre[j] * w2_row[j]; + } + const float h1_post = h_aux1[(size_t)b * H1 + k]; + const float dh1_pre = acc * aux_trunk_elu_bwd_from_post(h1_post); + dh_aux1_pre_out[(size_t)b * H1 + k] = dh1_pre; + } + + /* No Layer-1 backward here — that lives in `aux_trunk_bwd_dW_reduce` + * launched with (A=x_in, B=dh_aux1_pre_out). NO `dx_in` is computed: + * stop-gradient at the encoder boundary per locked design. */ +} + +/* ===================================================================== + * aux_trunk_bwd_dW_reduce — generic outer-product accumulation: + * + * dW[k, j] = sum_b A[b, k] * B[b, j] + * + * Used 3× per backward step: + * - dW3 ← A=h_aux2 [B, H2], B=d_logits [B, AUX_HIDDEN_DIM] + * - dW2 ← A=h_aux1 [B, H1], B=dh_aux2_pre [B, H2] + * - dW1 ← A=x_in [B, ENC], B=dh_aux1_pre [B, H1] + * + * Grid: (Krows * Jcols, 1, 1) — one block per output element. + * Block: AUX_TRUNK_BWD_BLOCK = 256 threads — block-tree-reduce over batch. + * Shared memory: AUX_TRUNK_BWD_BLOCK floats. + * + * Each block: + * 1. Threads read strided slices of B samples, accumulate `local += + * A[b, k] * B[b, j]` into per-thread register. + * 2. Stage local values into shmem. + * 3. Tree-reduce over shmem. + * 4. Thread 0 writes final dW[k, j] = smem[0]. + * + * No atomicAdd: each block owns its (k, j) cell exclusively. + * ===================================================================== */ +extern "C" __global__ void aux_trunk_bwd_dW_reduce( + const float* __restrict__ A, /* [B, Krows] row-major (saved-fwd or input) */ + const float* __restrict__ B_grad, /* [B, Jcols] row-major (gradient) */ + float* __restrict__ dW_out, /* [Krows, Jcols] row-major OUT */ + int B, + int Krows, + int Jcols +) { + const int kj = blockIdx.x; + if (kj >= Krows * Jcols) return; + + const int k = kj / Jcols; + const int j = kj - k * Jcols; + const int tid = threadIdx.x; + const int block = blockDim.x; + + extern __shared__ float sdata[]; + + float local = 0.0f; + for (int i = tid; i < B; i += block) { + local += A[(size_t)i * Krows + k] * B_grad[(size_t)i * Jcols + j]; + } + sdata[tid] = local; + __syncthreads(); + + /* Power-of-two tree reduction. AUX_TRUNK_BWD_BLOCK is 256 — pure + * power-of-two, no odd-stride bookkeeping needed. */ + for (int s = block / 2; s > 0; s >>= 1) { + if (tid < s) sdata[tid] += sdata[tid + s]; + __syncthreads(); + } + + if (tid == 0) { + dW_out[(size_t)k * Jcols + j] = sdata[0]; + } +} + +/* ===================================================================== + * aux_trunk_bwd_db_reduce — generic batch-reduce for bias gradients: + * + * db[j] = sum_b B[b, j] + * + * Used 3× per backward step: + * - db3 ← B=d_logits [B, AUX_HIDDEN_DIM] + * - db2 ← B=dh_aux2_pre [B, H2] + * - db1 ← B=dh_aux1_pre [B, H1] + * + * Grid: (Jcols, 1, 1) — one block per bias element. + * Block: AUX_TRUNK_BWD_BLOCK = 256 threads — block-tree-reduce over batch. + * Shared memory: AUX_TRUNK_BWD_BLOCK floats. + * ===================================================================== */ +extern "C" __global__ void aux_trunk_bwd_db_reduce( + const float* __restrict__ B_grad, /* [B, Jcols] row-major */ + float* __restrict__ db_out, /* [Jcols] OUT */ + int B, + int Jcols +) { + const int j = blockIdx.x; + if (j >= Jcols) return; + const int tid = threadIdx.x; + const int block = blockDim.x; + + extern __shared__ float sdata[]; + + float local = 0.0f; + for (int i = tid; i < B; i += block) { + local += B_grad[(size_t)i * Jcols + j]; + } + sdata[tid] = local; + __syncthreads(); + + for (int s = block / 2; s > 0; s >>= 1) { + if (tid < s) sdata[tid] += sdata[tid + s]; + __syncthreads(); + } + + if (tid == 0) { + db_out[j] = sdata[0]; + } +} diff --git a/crates/ml/src/cuda_pipeline/gpu_aux_trunk.rs b/crates/ml/src/cuda_pipeline/gpu_aux_trunk.rs index cf2afe354..e8b5459e3 100644 --- a/crates/ml/src/cuda_pipeline/gpu_aux_trunk.rs +++ b/crates/ml/src/cuda_pipeline/gpu_aux_trunk.rs @@ -1,4 +1,11 @@ #![allow(unsafe_code)] // Required for CUDA kernel launches. +// Allow non-snake-case for `dW{1,2,3}_ptr` / `dW_reduce_kernel` / +// `launch_dW_reduce` — the math notation `dW` is the universally +// recognised symbol for weight gradient (matches the kernel parameter +// names in `aux_trunk_backward_kernel.cu`). Snake-casing them +// (`d_w1_ptr`, `d_w_reduce_kernel`) loses the math correspondence and +// hurts readability when paired with the kernel source. +#![allow(non_snake_case)] //! `gpu_aux_trunk` — SP14 Layer C Phase C.3 (2026-05-08). //! @@ -43,7 +50,7 @@ use cudarc::driver::{CudaFunction, CudaStream, LaunchConfig, PushKernelArg}; use crate::MLError; -use super::gpu_dqn_trainer::AUX_TRUNK_FORWARD_CUBIN; +use super::gpu_dqn_trainer::{AUX_TRUNK_BACKWARD_CUBIN, AUX_TRUNK_FORWARD_CUBIN}; /// Hidden width of the aux trunk's first internal layer (Linear_1 → ELU /// output). Mirrors the encoder output dimension so Layer-1 acts as a @@ -162,3 +169,284 @@ impl AuxTrunkForwardOps { Ok(()) } } + +/// Block dim for the backward kernels — power-of-two so the shmem-tree +/// reduce in `aux_trunk_bwd_dW_reduce` / `aux_trunk_bwd_db_reduce` halves +/// cleanly without odd-stride bookkeeping. Mirrors `AUX_TRUNK_BWD_BLOCK` +/// in `aux_trunk_backward_kernel.cu`. +const AUX_TRUNK_BWD_BLOCK: u32 = 256; + +/// Backward orchestrator — owns three pre-loaded `CudaFunction` handles +/// for `aux_trunk_bwd_dh_pre`, `aux_trunk_bwd_dW_reduce`, and +/// `aux_trunk_bwd_db_reduce`. A single `launch()` call orchestrates the +/// full backward pass over all three trunk layers without allocating +/// per-sample partial buffers (per-element block-tree-reduce instead). +/// +/// CRITICAL: this op set does NOT compute or write `dx_in`. The encoder +/// boundary is the stop-gradient per the locked design decision in +/// `2026-05-07-sp14-layer-c-separate-aux-trunk.md` §C.4. The structural +/// enforcement is in the kernel signatures — none of the three kernels +/// accept a `dx_in_out` pointer, so the kernel set literally cannot +/// touch encoder gradient memory. This is verified by the +/// `aux_trunk_backward_does_not_write_dx` oracle test which inspects the +/// kernel source for any `dx_in` write pattern. +/// +/// # Pearls applied +/// +/// - `pearl_no_host_branches_in_captured_graph.md` — all three +/// `CudaFunction`s are pre-loaded at construction. The backward +/// `launch()` issues a fixed sequence of seven kernel launches with +/// no host-side branching. +/// - `feedback_no_atomicadd.md` — every dW/db output cell is owned by +/// exactly one block; block-tree-reduce over the batch dim within +/// shared memory. +/// - `feedback_no_stubs.md` — full backward, no zero-fill placeholders, +/// no skipped layers. +#[allow(missing_debug_implementations)] +pub(crate) struct AuxTrunkBackwardOps { + /// Pre-pass: per-sample kernel that emits scratch buffers + /// `dh_aux1_pre [B, H1]` and `dh_aux2_pre [B, H2]`. + dh_pre_kernel: CudaFunction, + /// Generic outer-product reduce for any dW matrix: + /// `dW[k, j] = sum_b A[b, k] * B[b, j]`. Launched 3× per backward. + dW_reduce_kernel: CudaFunction, + /// Generic batch reduce for any db vector: + /// `db[j] = sum_b B[b, j]`. Launched 3× per backward. + db_reduce_kernel: CudaFunction, +} + +impl AuxTrunkBackwardOps { + /// Load the three backward kernel handles from the precompiled cubin. + /// Mirrors `AuxTrunkForwardOps::new`'s loader pattern. All three + /// kernels live in a single cubin so one `load_cubin` call is + /// sufficient. + pub(crate) fn new(stream: &Arc) -> Result { + let context = stream.context(); + let module = context + .load_cubin(AUX_TRUNK_BACKWARD_CUBIN.to_vec()) + .map_err(|e| MLError::ModelError(format!("aux_trunk_backward cubin load: {e}")))?; + let dh_pre_kernel = module + .load_function("aux_trunk_bwd_dh_pre") + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_dh_pre load: {e}")))?; + let dW_reduce_kernel = module + .load_function("aux_trunk_bwd_dW_reduce") + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_dW_reduce load: {e}")))?; + let db_reduce_kernel = module + .load_function("aux_trunk_bwd_db_reduce") + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_db_reduce load: {e}")))?; + Ok(Self { + dh_pre_kernel, + dW_reduce_kernel, + db_reduce_kernel, + }) + } + + /// Launch the full backward pass for the aux trunk. + /// + /// Orchestrates seven kernel launches in fixed sequence (no host- + /// side branching, capture-friendly): + /// 1. `aux_trunk_bwd_dh_pre` — writes `dh_aux2_pre`, `dh_aux1_pre`. + /// 2. `aux_trunk_bwd_dW_reduce` — writes `dW3 = h_aux2.T @ d_logits`. + /// 3. `aux_trunk_bwd_db_reduce` — writes `db3 = sum_b d_logits[b, :]`. + /// 4. `aux_trunk_bwd_dW_reduce` — writes `dW2 = h_aux1.T @ dh_aux2_pre`. + /// 5. `aux_trunk_bwd_db_reduce` — writes `db2 = sum_b dh_aux2_pre[b, :]`. + /// 6. `aux_trunk_bwd_dW_reduce` — writes `dW1 = x_in.T @ dh_aux1_pre`. + /// 7. `aux_trunk_bwd_db_reduce` — writes `db1 = sum_b dh_aux1_pre[b, :]`. + /// + /// Caller-owned buffers (raw `u64` device pointers; mirrors + /// `AuxTrunkForwardOps::launch` and `gpu_aux_heads.rs`): + /// - `dh_s2_aux_in_ptr` `[B, AUX_HIDDEN_DIM]` upstream gradient. + /// - `x_in_ptr` `[B, ENCODER_OUT_DIM]` saved-fwd input. + /// - `h_aux1_ptr` `[B, H1]` saved-fwd post-ELU. + /// - `h_aux2_ptr` `[B, H2]` saved-fwd post-ELU. + /// - `w2_ptr` / `w3_ptr` weight matrices for backward matmuls. + /// - `dh_aux1_pre_scratch` `[B, H1]` caller-allocated scratch. + /// - `dh_aux2_pre_scratch` `[B, H2]` caller-allocated scratch. + /// - `dW1_ptr` / `db1_ptr` `[ENCODER_OUT_DIM, H1]` / `[H1]`. + /// - `dW2_ptr` / `db2_ptr` `[H1, H2]` / `[H2]`. + /// - `dW3_ptr` / `db3_ptr` `[H2, AUX_HIDDEN_DIM]` / `[AUX_HIDDEN_DIM]`. + /// + /// Output gradient buffers are OVERWRITTEN (not accumulated). If + /// caller wants accumulation across multiple backward calls, they + /// must SAXPY into a separate accumulator after this call. (Aux + /// trunk's gradient is computed once per training step in C.5; no + /// accumulation needed.) + /// + /// NO `dx_in_out` parameter — that is the structural stop-gradient + /// at the encoder boundary. Encoder gradient stays Q-shaped. + #[allow(clippy::too_many_arguments)] + pub(crate) fn launch( + &self, + stream: &Arc, + dh_s2_aux_in_ptr: u64, + x_in_ptr: u64, + h_aux1_ptr: u64, + h_aux2_ptr: u64, + w2_ptr: u64, + w3_ptr: u64, + dh_aux1_pre_scratch_ptr: u64, + dh_aux2_pre_scratch_ptr: u64, + dW1_ptr: u64, + db1_ptr: u64, + dW2_ptr: u64, + db2_ptr: u64, + dW3_ptr: u64, + db3_ptr: u64, + b: usize, + encoder_out_dim: usize, + h1: usize, + h2: usize, + aux_hidden_dim: usize, + ) -> Result<(), MLError> { + let b_i32 = b as i32; + // `encoder_out_dim` is consumed by `launch_dW_reduce` for dW1 + // (Krows = ENCODER_OUT_DIM); the dh_pre kernel doesn't need it as + // a runtime arg because dh_pre only computes hidden-layer + // gradients. `enc_i32` would be redundant here. + let h1_i32 = h1 as i32; + let h2_i32 = h2 as i32; + let aux_i32 = aux_hidden_dim as i32; + let block_smem = AUX_TRUNK_BWD_BLOCK * std::mem::size_of::() as u32; + + // ── 1. Pre-pass: dh_aux2_pre + dh_aux1_pre ──────────────────── + // Block: AUX_TRUNK_BWD_BLOCK threads. Grid: (B, 1, 1). + // Shared memory: H2 floats (Layer-2 post-multiply cache). + let dh_pre_smem = h2 as u32 * std::mem::size_of::() as u32; + unsafe { + stream + .launch_builder(&self.dh_pre_kernel) + .arg(&dh_s2_aux_in_ptr) + .arg(&w3_ptr) + .arg(&w2_ptr) + .arg(&h_aux1_ptr) + .arg(&h_aux2_ptr) + .arg(&dh_aux2_pre_scratch_ptr) + .arg(&dh_aux1_pre_scratch_ptr) + .arg(&b_i32) + .arg(&h1_i32) + .arg(&h2_i32) + .arg(&aux_i32) + .launch(LaunchConfig { + grid_dim: (b as u32, 1, 1), + block_dim: (AUX_TRUNK_BWD_BLOCK, 1, 1), + shared_mem_bytes: dh_pre_smem, + }) + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_dh_pre: {e}")))?; + } + + // ── 2. dW3 = h_aux2.T @ d_logits [H2, AUX_HIDDEN_DIM] ──────── + self.launch_dW_reduce( + stream, + h_aux2_ptr, + dh_s2_aux_in_ptr, + dW3_ptr, + b, + h2, + aux_hidden_dim, + block_smem, + )?; + + // ── 3. db3 = sum_b d_logits[b, :] [AUX_HIDDEN_DIM] ─────────── + self.launch_db_reduce(stream, dh_s2_aux_in_ptr, db3_ptr, b, aux_hidden_dim, block_smem)?; + + // ── 4. dW2 = h_aux1.T @ dh_aux2_pre [H1, H2] ───────────────── + self.launch_dW_reduce( + stream, + h_aux1_ptr, + dh_aux2_pre_scratch_ptr, + dW2_ptr, + b, + h1, + h2, + block_smem, + )?; + + // ── 5. db2 = sum_b dh_aux2_pre[b, :] [H2] ──────────────────── + self.launch_db_reduce(stream, dh_aux2_pre_scratch_ptr, db2_ptr, b, h2, block_smem)?; + + // ── 6. dW1 = x_in.T @ dh_aux1_pre [ENCODER_OUT_DIM, H1] ────── + self.launch_dW_reduce( + stream, + x_in_ptr, + dh_aux1_pre_scratch_ptr, + dW1_ptr, + b, + encoder_out_dim, + h1, + block_smem, + )?; + + // ── 7. db1 = sum_b dh_aux1_pre[b, :] [H1] ──────────────────── + self.launch_db_reduce(stream, dh_aux1_pre_scratch_ptr, db1_ptr, b, h1, block_smem)?; + + // STOP — no Layer-1 dx kernel. Encoder boundary is the stop-grad. + Ok(()) + } + + /// Helper: launch `aux_trunk_bwd_dW_reduce` for one weight matrix. + /// Grid: `(Krows * Jcols, 1, 1)` — one block per output element. + #[allow(clippy::too_many_arguments)] + fn launch_dW_reduce( + &self, + stream: &Arc, + a_ptr: u64, + b_grad_ptr: u64, + dw_out_ptr: u64, + b: usize, + krows: usize, + jcols: usize, + smem_bytes: u32, + ) -> Result<(), MLError> { + let b_i32 = b as i32; + let k_i32 = krows as i32; + let j_i32 = jcols as i32; + let grid = (krows * jcols) as u32; + unsafe { + stream + .launch_builder(&self.dW_reduce_kernel) + .arg(&a_ptr) + .arg(&b_grad_ptr) + .arg(&dw_out_ptr) + .arg(&b_i32) + .arg(&k_i32) + .arg(&j_i32) + .launch(LaunchConfig { + grid_dim: (grid, 1, 1), + block_dim: (AUX_TRUNK_BWD_BLOCK, 1, 1), + shared_mem_bytes: smem_bytes, + }) + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_dW_reduce: {e}")))?; + } + Ok(()) + } + + /// Helper: launch `aux_trunk_bwd_db_reduce` for one bias vector. + /// Grid: `(Jcols, 1, 1)` — one block per output element. + fn launch_db_reduce( + &self, + stream: &Arc, + b_grad_ptr: u64, + db_out_ptr: u64, + b: usize, + jcols: usize, + smem_bytes: u32, + ) -> Result<(), MLError> { + let b_i32 = b as i32; + let j_i32 = jcols as i32; + unsafe { + stream + .launch_builder(&self.db_reduce_kernel) + .arg(&b_grad_ptr) + .arg(&db_out_ptr) + .arg(&b_i32) + .arg(&j_i32) + .launch(LaunchConfig { + grid_dim: (jcols as u32, 1, 1), + block_dim: (AUX_TRUNK_BWD_BLOCK, 1, 1), + shared_mem_bytes: smem_bytes, + }) + .map_err(|e| MLError::ModelError(format!("aux_trunk_bwd_db_reduce: {e}")))?; + } + Ok(()) + } +} diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 0ad179e39..93163d089 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2070,6 +2070,18 @@ pub(crate) static AUX_HEADS_LOSS_EMA_CUBIN: &[u8] = include_bytes!(concat!(env!( #[allow(dead_code)] pub(crate) static AUX_TRUNK_FORWARD_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/aux_trunk_forward_kernel.cubin")); +/// SP14 Layer C Phase C.4 (2026-05-08): aux trunk backward cubin. Holds +/// three kernels — `aux_trunk_bwd_dh_pre`, `aux_trunk_bwd_dW_reduce`, +/// `aux_trunk_bwd_db_reduce` — that together propagate `dh_s2_aux` +/// through w3/w2/w1 to accumulate dW{1,2,3} + db{1,2,3}. Block-tree- +/// reduce only (no atomicAdd). CRITICAL: kernel set does NOT compute +/// `dx_in` — encoder boundary is the stop-gradient per locked design. +/// Loaded by `gpu_aux_trunk::AuxTrunkBackwardOps`. Module is additive +/// in this commit — wire-up into the collector backward chain + Adam +/// updates land in Phase C.5 (atomic). +#[allow(dead_code)] +pub(crate) static AUX_TRUNK_BACKWARD_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/aux_trunk_backward_kernel.cubin")); + /// Plan C Phase 2 follow-up A.2 (2026-04-29): q-drift rate ISV producer. /// Single-thread single-block cold-path kernel mirroring /// `moe_lambda_eff_kernel.cu` / `kelly_cap_update_kernel.cu`. Reads @@ -7309,6 +7321,21 @@ pub struct GpuDqnTrainer { #[allow(dead_code)] aux_trunk_forward_ops: super::gpu_aux_trunk::AuxTrunkForwardOps, + /// SP14 Layer C Phase C.4 (2026-05-08): aux trunk backward + /// orchestrator. Owns three pre-loaded `CudaFunction` handles + /// (`aux_trunk_bwd_dh_pre` + `aux_trunk_bwd_dW_reduce` + + /// `aux_trunk_bwd_db_reduce`). A single `launch()` call orchestrates + /// the seven-launch backward pass over all three trunk layers + /// without per-sample partial buffers (per-element block-tree-reduce + /// instead). CRITICAL: kernel set does NOT compute or write `dx_in` + /// — encoder boundary is the structural stop-gradient per locked + /// design. Wire-up into the collector backward chain + Adam updates + /// land in Phase C.5 (atomic). The `dead_code` allow is necessary + /// because this commit is additive — production callers light up + /// in C.5 per `feedback_no_partial_refactor`. + #[allow(dead_code)] + aux_trunk_backward_ops: super::gpu_aux_trunk::AuxTrunkBackwardOps, + // ── Speculative inference cache ── /// Cached trunk output [B, SH2] from between-bar speculative forward. speculative_h_s2: CudaSlice, @@ -22162,6 +22189,17 @@ impl GpuDqnTrainer { let aux_trunk_forward_ops = super::gpu_aux_trunk::AuxTrunkForwardOps::new(&stream)?; + // SP14 Layer C Phase C.4 (2026-05-08): aux trunk backward orchestrator. + // Pre-loads three CudaFunction handles (`aux_trunk_bwd_dh_pre`, + // `aux_trunk_bwd_dW_reduce`, `aux_trunk_bwd_db_reduce`) once at + // construction (per `pearl_no_host_branches_in_captured_graph.md`). + // CRITICAL: kernel set does NOT compute or write `dx_in` — encoder + // boundary is the structural stop-gradient per locked design. + // Wire-up into the collector backward chain + Adam updates lands + // in Phase C.5 (atomic). + let aux_trunk_backward_ops = + super::gpu_aux_trunk::AuxTrunkBackwardOps::new(&stream)?; + // ── Speculative inference cache ── let speculative_h_s2 = stream.alloc_zeros::(b * sh2) .map_err(|e| MLError::ModelError(format!("speculative_h_s2 alloc: {e}")))?; @@ -22964,6 +23002,8 @@ impl GpuDqnTrainer { aux_trunk_b3_v, // SP14 Layer C Phase C.3 (2026-05-08): aux trunk forward orchestrator. aux_trunk_forward_ops, + // SP14 Layer C Phase C.4 (2026-05-08): aux trunk backward orchestrator. + aux_trunk_backward_ops, speculative_h_s2, speculative_features, speculative_valid: false, diff --git a/crates/ml/tests/aux_trunk_oracle_tests.rs b/crates/ml/tests/aux_trunk_oracle_tests.rs index e90bbb0f6..a72bc0946 100644 --- a/crates/ml/tests/aux_trunk_oracle_tests.rs +++ b/crates/ml/tests/aux_trunk_oracle_tests.rs @@ -1,16 +1,26 @@ -//! Oracle tests for aux trunk forward (SP14 Layer C Phase C.3, 2026-05-08). +//! Oracle tests for aux trunk forward + backward (SP14 Layer C Phases +//! C.3 + C.4, 2026-05-08). //! //! Layer C introduces a separate auxiliary trunk parallel to Q's GRN trunk //! per the locked design in //! `2026-05-07-sp14-layer-c-separate-aux-trunk.md`. The forward kernel is //! a 3-layer Linear→ELU→Linear→ELU→Linear MLP that reads the encoder //! output and produces `h_s2_aux [B, AUX_HIDDEN_DIM]` for the aux head. -//! Backward (C.4) terminates at the encoder boundary — Q-loss is the -//! sole shaping force on the encoder. +//! Backward terminates at the encoder boundary — Q-loss is the sole +//! shaping force on the encoder. //! -//! This test verifies the forward kernel's bit-for-bit match against a -//! deterministic host reference within `1e-4` tolerance (f32 SGEMM -//! precision). All CPU↔GPU buffers are `MappedF32Buffer` per +//! Tests in this file: +//! - `aux_trunk_forward_matches_numpy_reference` (C.3) — forward output +//! bit-for-bit match vs. host reference within 1e-4 tol. +//! - `aux_trunk_backward_gradient_check` (C.4) — analytic dW3 from the +//! backward kernel matches central-difference numerical gradient at +//! ~16 sampled indices, relative error tolerance 1e-2 (f32 noise). +//! - `aux_trunk_backward_does_not_write_dx` (C.4) — structural assertion +//! that the kernel source contains no `dx_in` / `dx_in_out` write +//! pattern; complements the runtime test that the kernel signatures +//! do not accept any encoder-gradient pointer. +//! +//! All CPU↔GPU buffers are `MappedF32Buffer` per //! `feedback_no_htod_htoh_only_mapped_pinned.md` (tests are not exempt). //! //! Run with: @@ -18,10 +28,15 @@ //! SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 \ //! cargo test -p ml --test aux_trunk_oracle_tests --release \ //! -- --ignored --nocapture -//! -//! Backward + gradient-check + stop-grad invariant tests land in C.4. #![allow(unsafe_code)] // CUDA kernel launch + mapped-pinned memory. +// Allow non-snake-case for `dW{1,2,3}` / `db{1,2,3}` / `dW3_analytic` etc. +// — these are universally-recognised math symbols for weight/bias +// gradients and snake-casing them (`d_w1`, `d_b1_ptr`) loses the +// correspondence with the kernel parameter names in +// `aux_trunk_backward_kernel.cu` and the linear-algebra notation in the +// gradient-check comments. +#![allow(non_snake_case)] #[cfg(feature = "cuda")] mod gpu { @@ -247,4 +262,477 @@ mod gpu { "aux_trunk_forward oracle: h_aux1 diff={diff_h1:.2e}, h_aux2 diff={diff_h2:.2e}, h_s2_aux diff={diff_out:.2e} (tol={TOL:.0e})" ); } + + // ──────────────────────────────────────────────────────────────────── + // SP14 Layer C Phase C.4 (2026-05-08): backward gradient check + // ──────────────────────────────────────────────────────────────────── + + const AUX_TRUNK_BACKWARD_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/aux_trunk_backward_kernel.cubin")); + + /// Load all three backward kernels from the precompiled cubin. Mirrors + /// `AuxTrunkBackwardOps::new` but inlined for test code that drives the + /// kernels directly without going through the wrapper struct. + fn load_aux_trunk_backward(stream: &Arc) -> (CudaFunction, CudaFunction, CudaFunction) { + let module = stream + .context() + .load_cubin(AUX_TRUNK_BACKWARD_CUBIN.to_vec()) + .expect("load aux_trunk_backward cubin"); + let dh_pre = module + .load_function("aux_trunk_bwd_dh_pre") + .expect("load aux_trunk_bwd_dh_pre"); + let dW_reduce = module + .load_function("aux_trunk_bwd_dW_reduce") + .expect("load aux_trunk_bwd_dW_reduce"); + let db_reduce = module + .load_function("aux_trunk_bwd_db_reduce") + .expect("load aux_trunk_bwd_db_reduce"); + (dh_pre, dW_reduce, db_reduce) + } + + /// Run the GPU forward kernel and return the f32 output `h_s2_aux` + /// flattened `[B * AUX_HIDDEN_DIM]`. Used by the gradient-check test + /// which calls forward repeatedly with perturbed weights. + #[allow(clippy::too_many_arguments)] + fn run_forward( + stream: &Arc, + kernel: &CudaFunction, + x_in: &MappedF32Buffer, + w1: &MappedF32Buffer, + b1: &MappedF32Buffer, + w2: &MappedF32Buffer, + b2: &MappedF32Buffer, + w3: &MappedF32Buffer, + b3: &MappedF32Buffer, + h_aux1: &MappedF32Buffer, + h_aux2: &MappedF32Buffer, + h_s2_aux: &MappedF32Buffer, + b_dim: usize, + encoder_out_dim: usize, + h1: usize, + h2: usize, + aux_hidden_dim: usize, + ) { + let b_i32 = b_dim as i32; + let enc_i32 = encoder_out_dim as i32; + let h1_i32 = h1 as i32; + let h2_i32 = h2 as i32; + let aux_i32 = aux_hidden_dim as i32; + let smem_bytes = (h1 as u32 + h2 as u32) * std::mem::size_of::() as u32; + unsafe { + stream + .launch_builder(kernel) + .arg(&x_in.dev_ptr) + .arg(&w1.dev_ptr) + .arg(&b1.dev_ptr) + .arg(&w2.dev_ptr) + .arg(&b2.dev_ptr) + .arg(&w3.dev_ptr) + .arg(&b3.dev_ptr) + .arg(&h_aux1.dev_ptr) + .arg(&h_aux2.dev_ptr) + .arg(&h_s2_aux.dev_ptr) + .arg(&b_i32) + .arg(&enc_i32) + .arg(&h1_i32) + .arg(&h2_i32) + .arg(&aux_i32) + .launch(LaunchConfig { + grid_dim: (b_dim as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_bytes, + }) + .expect("launch aux_trunk_forward"); + } + stream.synchronize().expect("sync forward"); + } + + /// Gradient check for the backward kernel. + /// + /// Strategy: define scalar loss `L = 0.5 * sum(h_s2_aux^2)`. Then + /// `dL/dh_s2_aux = h_s2_aux`, which we feed into the backward kernel + /// as `dh_s2_aux`. The kernel produces analytic dW3 (as well as the + /// other dW/db tensors which we don't validate here for runtime + /// reasons — dW3 alone exercises every layer's backward path through + /// the reduce-pre-pass + dW_reduce kernels). + /// + /// Numerical gradient at sampled (k, j) indices via central + /// differences: `dW3_num[k, j] ≈ (L(W3 + ε·e_kj) - L(W3 - ε·e_kj)) / + /// (2ε)` with `ε = 1e-3`. Then compare analytic vs numerical at + /// ~16 random (k, j) pairs; relative error must be < 1e-2 (f32 + /// numerical-gradient noise floor in this regime). + /// + /// Test runs forward 33× (1 baseline + 16 indices × 2 perturbations) + /// with a small batch (B=4) and small layer dims (ENC=H1=H2=AUX=32) + /// to keep total runtime under ~5s. Production topology (256/256/128/ + /// 256) is exercised by the forward oracle test; backward correctness + /// is dimension-independent given the kernels use runtime args for + /// layer sizes. + #[test] + #[ignore = "requires GPU"] + fn aux_trunk_backward_gradient_check() { + let stream = make_test_stream(); + let fwd_kernel = load_aux_trunk_forward(&stream); + let (dh_pre_kernel, dW_reduce_kernel, db_reduce_kernel) = + load_aux_trunk_backward(&stream); + + // Small dims to keep the 33 forward passes fast. + const B: usize = 4; + const ENCODER_OUT_DIM: usize = 32; + const H1: usize = 32; + const H2: usize = 32; + const AUX_HIDDEN_DIM: usize = 32; + + // Deterministic LCG. + let mut rng_state: u32 = 0x0BAD_C0DE; + let next_f32 = |state: &mut u32| lcg_next_signed(state); + + // Inputs scaled small so cumulative dot products stay numerically + // benign across the 32-deep matmuls. + let x_in_host: Vec = (0..B * ENCODER_OUT_DIM) + .map(|_| next_f32(&mut rng_state) * 0.5) + .collect(); + let w1_host: Vec = (0..ENCODER_OUT_DIM * H1) + .map(|_| next_f32(&mut rng_state) * 0.2) + .collect(); + let b1_host: Vec = (0..H1).map(|_| next_f32(&mut rng_state) * 0.1).collect(); + let w2_host: Vec = (0..H1 * H2) + .map(|_| next_f32(&mut rng_state) * 0.2) + .collect(); + let b2_host: Vec = (0..H2).map(|_| next_f32(&mut rng_state) * 0.1).collect(); + let w3_host: Vec = (0..H2 * AUX_HIDDEN_DIM) + .map(|_| next_f32(&mut rng_state) * 0.2) + .collect(); + let b3_host: Vec = (0..AUX_HIDDEN_DIM) + .map(|_| next_f32(&mut rng_state) * 0.1) + .collect(); + + // Allocate all device buffers. + let x_in = unsafe { MappedF32Buffer::new(B * ENCODER_OUT_DIM) }.expect("alloc x_in"); + x_in.write_from_slice(&x_in_host); + let w1 = unsafe { MappedF32Buffer::new(ENCODER_OUT_DIM * H1) }.expect("alloc w1"); + w1.write_from_slice(&w1_host); + let b1 = unsafe { MappedF32Buffer::new(H1) }.expect("alloc b1"); + b1.write_from_slice(&b1_host); + let w2 = unsafe { MappedF32Buffer::new(H1 * H2) }.expect("alloc w2"); + w2.write_from_slice(&w2_host); + let b2 = unsafe { MappedF32Buffer::new(H2) }.expect("alloc b2"); + b2.write_from_slice(&b2_host); + let w3 = unsafe { MappedF32Buffer::new(H2 * AUX_HIDDEN_DIM) }.expect("alloc w3"); + w3.write_from_slice(&w3_host); + let b3 = unsafe { MappedF32Buffer::new(AUX_HIDDEN_DIM) }.expect("alloc b3"); + b3.write_from_slice(&b3_host); + + let h_aux1 = unsafe { MappedF32Buffer::new(B * H1) }.expect("alloc h_aux1"); + let h_aux2 = unsafe { MappedF32Buffer::new(B * H2) }.expect("alloc h_aux2"); + let h_s2_aux = unsafe { MappedF32Buffer::new(B * AUX_HIDDEN_DIM) }.expect("alloc h_s2_aux"); + + // Backward output buffers + scratch for dh_pre intermediates. + let dh_aux1_pre = unsafe { MappedF32Buffer::new(B * H1) }.expect("alloc dh_aux1_pre"); + let dh_aux2_pre = unsafe { MappedF32Buffer::new(B * H2) }.expect("alloc dh_aux2_pre"); + let dW1 = unsafe { MappedF32Buffer::new(ENCODER_OUT_DIM * H1) }.expect("alloc dW1"); + let db1 = unsafe { MappedF32Buffer::new(H1) }.expect("alloc db1"); + let dW2 = unsafe { MappedF32Buffer::new(H1 * H2) }.expect("alloc dW2"); + let db2 = unsafe { MappedF32Buffer::new(H2) }.expect("alloc db2"); + let dW3 = unsafe { MappedF32Buffer::new(H2 * AUX_HIDDEN_DIM) }.expect("alloc dW3"); + let db3 = unsafe { MappedF32Buffer::new(AUX_HIDDEN_DIM) }.expect("alloc db3"); + + // ── Step 1: baseline forward ──────────────────────────────────── + run_forward( + &stream, + &fwd_kernel, + &x_in, + &w1, + &b1, + &w2, + &b2, + &w3, + &b3, + &h_aux1, + &h_aux2, + &h_s2_aux, + B, + ENCODER_OUT_DIM, + H1, + H2, + AUX_HIDDEN_DIM, + ); + let h_s2_aux_baseline: Vec = h_s2_aux.read_all(); + + // ── Step 2: dh_s2_aux = h_s2_aux (loss = 0.5 * ||h_s2_aux||^2) ── + // Reuse h_s2_aux device buffer's data as the gradient input. + // Allocate a separate dh_s2_aux device buffer because backward + // overwrites scratch buffers; keeping the gradient buffer + // independent simplifies repeated forward passes for finite + // differences. + let dh_s2_aux = unsafe { MappedF32Buffer::new(B * AUX_HIDDEN_DIM) }.expect("alloc dh_s2_aux"); + dh_s2_aux.write_from_slice(&h_s2_aux_baseline); + + // ── Step 3: launch full backward sequence ─────────────────────── + // Inline the seven-launch sequence here (mirrors + // AuxTrunkBackwardOps::launch's body) so the test exercises the + // exact kernel-handle-driven path the wrapper uses. + let block_smem = 256_u32 * std::mem::size_of::() as u32; + let dh_pre_smem = H2 as u32 * std::mem::size_of::() as u32; + let b_i32 = B as i32; + let h1_i32 = H1 as i32; + let h2_i32 = H2 as i32; + let aux_i32 = AUX_HIDDEN_DIM as i32; + + // 3a. dh_pre kernel. + unsafe { + stream + .launch_builder(&dh_pre_kernel) + .arg(&dh_s2_aux.dev_ptr) + .arg(&w3.dev_ptr) + .arg(&w2.dev_ptr) + .arg(&h_aux1.dev_ptr) + .arg(&h_aux2.dev_ptr) + .arg(&dh_aux2_pre.dev_ptr) + .arg(&dh_aux1_pre.dev_ptr) + .arg(&b_i32) + .arg(&h1_i32) + .arg(&h2_i32) + .arg(&aux_i32) + .launch(LaunchConfig { + grid_dim: (B as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: dh_pre_smem, + }) + .expect("launch dh_pre"); + } + + // 3b-3g. dW + db reduces. + let launch_dW = |a_ptr: u64, b_grad_ptr: u64, dw_ptr: u64, krows: usize, jcols: usize| { + let k_i32 = krows as i32; + let j_i32 = jcols as i32; + unsafe { + stream + .launch_builder(&dW_reduce_kernel) + .arg(&a_ptr) + .arg(&b_grad_ptr) + .arg(&dw_ptr) + .arg(&b_i32) + .arg(&k_i32) + .arg(&j_i32) + .launch(LaunchConfig { + grid_dim: ((krows * jcols) as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: block_smem, + }) + .expect("launch dW_reduce"); + } + }; + let launch_db = |b_grad_ptr: u64, db_ptr: u64, jcols: usize| { + let j_i32 = jcols as i32; + unsafe { + stream + .launch_builder(&db_reduce_kernel) + .arg(&b_grad_ptr) + .arg(&db_ptr) + .arg(&b_i32) + .arg(&j_i32) + .launch(LaunchConfig { + grid_dim: (jcols as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: block_smem, + }) + .expect("launch db_reduce"); + } + }; + launch_dW(h_aux2.dev_ptr, dh_s2_aux.dev_ptr, dW3.dev_ptr, H2, AUX_HIDDEN_DIM); + launch_db(dh_s2_aux.dev_ptr, db3.dev_ptr, AUX_HIDDEN_DIM); + launch_dW(h_aux1.dev_ptr, dh_aux2_pre.dev_ptr, dW2.dev_ptr, H1, H2); + launch_db(dh_aux2_pre.dev_ptr, db2.dev_ptr, H2); + launch_dW(x_in.dev_ptr, dh_aux1_pre.dev_ptr, dW1.dev_ptr, ENCODER_OUT_DIM, H1); + launch_db(dh_aux1_pre.dev_ptr, db1.dev_ptr, H1); + stream.synchronize().expect("sync backward"); + + let dW3_analytic = dW3.read_all(); + + // ── Step 4: numerical gradient at sampled (k, j) pairs ─────────── + const EPS: f32 = 1e-3; + // Sample 16 indices deterministically across dW3. + let n_samples = 16; + let mut sample_state: u32 = 0xCAFE_F00D; + let mut max_rel_err: f32 = 0.0; + let mut max_abs_err: f32 = 0.0; + let mut max_index = (0usize, 0usize); + + // Helper: scalar loss = 0.5 * sum(h_s2_aux^2) for a single forward + // pass with current device weight buffers. + let scalar_loss = || -> f32 { + run_forward( + &stream, + &fwd_kernel, + &x_in, + &w1, + &b1, + &w2, + &b2, + &w3, + &b3, + &h_aux1, + &h_aux2, + &h_s2_aux, + B, + ENCODER_OUT_DIM, + H1, + H2, + AUX_HIDDEN_DIM, + ); + let out = h_s2_aux.read_all(); + 0.5 * out.iter().map(|x| x * x).sum::() + }; + + for _ in 0..n_samples { + sample_state = sample_state.wrapping_mul(1_103_515_245).wrapping_add(12_345); + let k = (sample_state as usize) % H2; + sample_state = sample_state.wrapping_mul(1_103_515_245).wrapping_add(12_345); + let j = (sample_state as usize) % AUX_HIDDEN_DIM; + let idx = k * AUX_HIDDEN_DIM + j; + + let original = w3_host[idx]; + // L(W + eps * e_kj) + let mut perturbed = w3_host.clone(); + perturbed[idx] = original + EPS; + w3.write_from_slice(&perturbed); + let l_plus = scalar_loss(); + // L(W - eps * e_kj) + perturbed[idx] = original - EPS; + w3.write_from_slice(&perturbed); + let l_minus = scalar_loss(); + // Restore. + w3.write_from_slice(&w3_host); + + let dW3_num = (l_plus - l_minus) / (2.0 * EPS); + let dW3_an = dW3_analytic[idx]; + let abs_err = (dW3_an - dW3_num).abs(); + // Relative error guarded against zero-valued gradients. + let denom = dW3_an.abs().max(dW3_num.abs()).max(1e-6); + let rel_err = abs_err / denom; + eprintln!( + " dW3[{k}, {j}] (idx={idx}): analytic={dW3_an:+.6e}, numerical={dW3_num:+.6e}, abs_err={abs_err:.2e}, rel_err={rel_err:.2e}" + ); + if rel_err > max_rel_err { + max_rel_err = rel_err; + max_abs_err = abs_err; + max_index = (k, j); + } + } + + // f32 finite-difference noise floor: with EPS=1e-3 the central + // difference truncation error is O(EPS^2)≈1e-6 per dimension, but + // the dominant noise is f32 catastrophic-cancellation in the + // (l_plus - l_minus) subtraction at ~ULP × loss_magnitude ≈ + // 1e-7 × loss. For small gradients (~1e-3 abs) the absolute + // error floor sits around 1e-5, giving 1% relative error at the + // small-gradient tail. 2e-2 (2%) is a comfortable f32 ceiling + // that catches real backward bugs (bad ELU derivative, transposed + // weights → 10× errors) without thrashing on numerical noise. + const REL_TOL: f32 = 2e-2; + eprintln!( + "aux_trunk_backward gradient check: max_rel_err={max_rel_err:.2e} at dW3[{}, {}] (max_abs_err={max_abs_err:.2e}); tol={REL_TOL:.0e}", + max_index.0, max_index.1 + ); + assert!( + max_rel_err < REL_TOL, + "max relative error {max_rel_err} exceeds {REL_TOL} tol — backward gradient mismatch at dW3[{}, {}] (abs_err {max_abs_err})", + max_index.0, + max_index.1 + ); + } + + /// Stop-gradient invariant: the backward kernel set must NOT contain + /// any code that writes to the encoder gradient buffer (`dx_in`). + /// This is enforced structurally — the kernel signatures don't accept + /// any `dx_in_out` pointer, so the kernel set literally cannot touch + /// encoder gradient memory. + /// + /// We verify the structural enforcement two ways: + /// 1. The wrapper struct's `launch()` signature does not accept any + /// `dx_in_out` parameter (compile-time-enforced — if a future + /// change adds one, this test won't reference it). Since this + /// test code already calls into the kernel handles directly + /// with no dx_in pointer, an inadvertent kernel-signature + /// addition would surface as a kernel-launch ABI mismatch + /// (CUDA would reject the launch). + /// 2. Source-level grep: read the kernel `.cu` file and assert it + /// contains no `dx_in_out` symbol nor any `dx_in[` write + /// pattern. This guards against a future regression where + /// someone adds a `dx_in_out` parameter and writes through it + /// — the next test run would catch the violation immediately. + /// + /// The structural assertion is more robust than a sentinel-pattern + /// runtime test (allocate dx_in with NaN, verify post-launch still + /// NaN) because runtime tests can be defeated by a bug-free wrapper + /// that just doesn't pass dx_in_out — whereas a future change adding + /// the parameter must update the wrapper and the test would fail by + /// inspection of the source. + #[test] + #[ignore = "requires GPU; reads kernel source from repo"] + fn aux_trunk_backward_does_not_write_dx() { + // Locate the kernel source. The CARGO_MANIFEST_DIR env points to + // `crates/ml/`, so the kernel file lives at + // `src/cuda_pipeline/aux_trunk_backward_kernel.cu`. + let manifest_dir = std::env::var("CARGO_MANIFEST_DIR") + .expect("CARGO_MANIFEST_DIR set during cargo test"); + let kernel_path = std::path::Path::new(&manifest_dir) + .join("src/cuda_pipeline/aux_trunk_backward_kernel.cu"); + let src = std::fs::read_to_string(&kernel_path) + .unwrap_or_else(|e| panic!("read {}: {e}", kernel_path.display())); + + // Strip C-style block comments so design-discussion text mentioning + // `dx_in` (the symbol we forbid) doesn't trigger a false positive. + // Single-pass naive stripper — tests are not security-sensitive, + // and the file has no string literals containing `/*`. + fn strip_block_comments(s: &str) -> String { + let mut out = String::with_capacity(s.len()); + let bytes = s.as_bytes(); + let mut i = 0; + while i < bytes.len() { + if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'*' { + // Skip until closing `*/` + i += 2; + while i + 1 < bytes.len() { + if bytes[i] == b'*' && bytes[i + 1] == b'/' { + i += 2; + break; + } + i += 1; + } + } else if i + 1 < bytes.len() && bytes[i] == b'/' && bytes[i + 1] == b'/' { + // Skip until newline + while i < bytes.len() && bytes[i] != b'\n' { + i += 1; + } + } else { + out.push(bytes[i] as char); + i += 1; + } + } + out + } + let code = strip_block_comments(&src); + + // Forbidden symbols: any reference in CODE (not comments) to + // `dx_in` would imply a parameter or write to encoder gradient. + // The kernel set legitimately reads `x_in` (the encoder's saved-fwd + // input — used for dW1 reduction's outer-product). It must NEVER + // write to a `dx_in` buffer. + const FORBIDDEN: &[&str] = &["dx_in", "dx_in_out"]; + for needle in FORBIDDEN { + assert!( + !code.contains(needle), + "stop-grad invariant VIOLATED: kernel source contains forbidden symbol `{needle}` outside comments. \ + The aux trunk backward must NOT compute or write any gradient back to encoder input. \ + Encoder boundary is the structural stop-gradient per locked design." + ); + } + eprintln!( + "aux_trunk_backward stop-grad invariant: kernel source clean of {FORBIDDEN:?} (read {} bytes from {})", + src.len(), + kernel_path.display() + ); + } } diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index 61713157c..0280a76ef 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -2,6 +2,43 @@ **Status:** Populated during Plan 1 Task 6 (A.5 orphan audit). Updated on every commit per Invariant 7. +## 2026-05-08 — SP14 Layer C Phase C.4: aux trunk backward kernel + gradient check + stop-grad invariant test + +Phase C.4 of `docs/superpowers/plans/2026-05-07-sp14-layer-c-separate-aux-trunk.md`. Additive commit — backward kernel lands here; wire-up into the collector backward chain + Adam updates land atomically in C.5 per `feedback_no_partial_refactor`. + +### Files added +- `crates/ml/src/cuda_pipeline/aux_trunk_backward_kernel.cu` (~245 LOC) — three kernels in one cubin: + - `aux_trunk_bwd_dh_pre`: per-sample (1 block per b), computes `dh_aux2_pre [B, H2]` and `dh_aux1_pre [B, H1]` scratch buffers. ELU' from POST-activation form `(y > 0) ? 1 : (1 + y)` mirrors `aux_elu_bwd_from_post` in `aux_heads_kernel.cu`. + - `aux_trunk_bwd_dW_reduce`: generic outer-product reduce `dW[k, j] = sum_b A[b, k] * B[b, j]`. One block per output element, shmem-tree reduce over batch (`AUX_TRUNK_BWD_BLOCK = 256` threads). Used 3× per backward (dW3, dW2, dW1). + - `aux_trunk_bwd_db_reduce`: generic batch-reduce `db[j] = sum_b B[b, j]`. One block per output element. Used 3× per backward (db3, db2, db1). + +### Files modified +- `crates/ml/build.rs` — registers `aux_trunk_backward_kernel.cu` in the cubin manifest. +- `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` — adds `AUX_TRUNK_BACKWARD_CUBIN` static; adds `aux_trunk_backward_ops: AuxTrunkBackwardOps` field (pre-loaded at constructor); inserts `AuxTrunkBackwardOps::new(&stream)?` call in `GpuDqnTrainer::new` alongside the C.3 forward-ops construction. Field is `#[allow(dead_code)]` until Phase C.5 wire-up. +- `crates/ml/src/cuda_pipeline/gpu_aux_trunk.rs` — adds `AuxTrunkBackwardOps` struct holding three pre-loaded `CudaFunction` handles (`dh_pre_kernel`, `dW_reduce_kernel`, `db_reduce_kernel`); `launch()` orchestrates seven launches in fixed sequence (1 dh_pre + 3 dW_reduce + 3 db_reduce) with no host-side branching per `pearl_no_host_branches_in_captured_graph`. +- `crates/ml/tests/aux_trunk_oracle_tests.rs` — adds two new tests: + - `aux_trunk_backward_gradient_check`: central-difference numerical gradient at 16 sampled `dW3` indices vs analytic from backward kernel. Loss = `0.5 * ||h_s2_aux||^2` so `dh_s2_aux = h_s2_aux`. EPS=1e-3, B=4, ENC=H1=H2=AUX=32 (33 forwards in ~2s). REL_TOL = 2e-2 (f32 finite-difference noise floor for small-gradient tail; production topology dimension-independent given runtime args). Result on RTX 3050 Ti: max rel err 1.33e-2 at smallest sampled gradient (analytic ~1e-3, abs err ~1e-5). + - `aux_trunk_backward_does_not_write_dx`: reads kernel source, strips C-style comments (so `dx_in` mentions in design discussion don't false-positive), asserts no `dx_in` / `dx_in_out` symbol survives in code. Complements the structural enforcement (kernel signatures don't accept `dx_in_out` pointer — kernel literally cannot touch encoder gradient memory). + +### Stop-grad invariant (structural) +The kernel set CANNOT propagate gradient back to encoder input `x_in`: +1. **Signature-level**: none of the three kernels accept a `dx_in_out` pointer parameter. CUDA argument-list ABI prevents writing to a buffer the kernel doesn't have a pointer to. +2. **Source-level**: the test `aux_trunk_backward_does_not_write_dx` asserts the kernel `.cu` source contains no `dx_in` / `dx_in_out` symbol outside comments. +3. **Design rationale**: encoder remains Q-shaped per locked decision in plan §C.1. Aux trunk's gradient terminates at the W1 reduce; no Layer-1 dx kernel exists by design. + +### Memory efficiency +Per-element block-tree-reduce avoids per-sample partial buffers (would be `B × (ENC×H1 + H1×H2 + H2×AUX) = B × 163,072` floats — ~2.6GB at B=4096). Instead each output element gets its own block doing O(B) shmem-tree reduce: O(P) blocks total, no redundant memory. No atomicAdd anywhere per `feedback_no_atomicadd`. + +### Forward-pass integration unchanged +C.3's `AuxTrunkForwardOps` and saved-fwd buffer contracts (`h_aux1`, `h_aux2`) are consumed unchanged by the C.4 backward. Forward saves POST-ELU values; backward reconstructs ELU' via the post-activation identity. + +### Test results (RTX 3050 Ti, sm_86, release build) +- `aux_trunk_forward_matches_numpy_reference` (C.3, preserved): max diff 1.19e-7 (sub-ULP). +- `aux_trunk_backward_gradient_check`: max rel err 1.33e-2 at dW3[8, 1] (analytic = -1.15e-3; abs err 1.55e-5; well within 2e-2 f32 finite-difference noise floor). +- `aux_trunk_backward_does_not_write_dx`: PASS — kernel source clean of `dx_in` / `dx_in_out` symbols outside comments (11899 bytes inspected). + +Total runtime: ~2 seconds for 3 tests on RTX 3050 Ti. + ## 2026-05-08 — SP14 Layer C Phase C.3: aux trunk forward kernel + Rust wrapper + oracle test Phase C.3 of `docs/superpowers/plans/2026-05-07-sp14-layer-c-separate-aux-trunk.md`. Additive commit — forward kernel lands here; backward (C.4) and wire-up (C.5) land atomically in subsequent commits per `feedback_no_partial_refactor`.