From c508d101c161ef394585da05b09fe587e93465ac Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 17 May 2026 23:27:20 +0200 Subject: [PATCH] plan(perf): K-loop block-per-batch parallelization implementation plan Bite-sized 22-task plan implementing the design at docs/superpowers/specs/2026-05-17-kloop-parallelization-design.md. Four atomic commits per the spec's Rollout section: Commit 1: reduce_axis0 kernel + cfc_step refactor (Tasks 1-10) Commit 2: GRN bwd refactor (Tasks 11-14) Commit 3: VSN bwd refactor (Tasks 15-17) Commit 4: attention_pool bwd refactor (Tasks 18-20) Each commit covers kernel rewrite + per-batch scratch buffers + reducer launches + memset_zeros wiring + smoke validation. Tests added across commits: - cfc_bwd_b1_oracle.rs (Task 6): B=1 oracle vs single-sample helper within relative_eq 1e-5 (FP-tolerant, not bit-exact) - stacked_trainer_loss_shrinks_at_batch_32 (Task 7): first test that exercises the cross-batch reducer path - cfc_bwd_scratch_clears_between_steps (Task 8): scratch-zero regression guard Acceptance gates 1-8 from the spec are mapped to Tasks 6, 7, 8, 9, 21, 22 (local + cluster A/B perf). Reference baseline for gate #8 is t6z89's per-epoch AUC trajectory. Co-Authored-By: Claude Opus 4.7 --- .../plans/2026-05-17-kloop-parallelization.md | 1861 +++++++++++++++++ 1 file changed, 1861 insertions(+) create mode 100644 docs/superpowers/plans/2026-05-17-kloop-parallelization.md diff --git a/docs/superpowers/plans/2026-05-17-kloop-parallelization.md b/docs/superpowers/plans/2026-05-17-kloop-parallelization.md new file mode 100644 index 000000000..5d159cadf --- /dev/null +++ b/docs/superpowers/plans/2026-05-17-kloop-parallelization.md @@ -0,0 +1,1861 @@ +# K-loop Parallelization Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Parallelize five backward / K-loop CUDA kernels across the batch dimension to remove the single-SM bottleneck observed in `ml-alpha` training (t6z89 at 8 min/epoch on L40S). + +**Architecture:** Block-per-batch refactor of `cfc_step_batched` (fwd+bwd), `multi_horizon_heads_grn_bwd_batched`, `variable_selection_bwd`, and `attention_pool_bwd`. Param-grad accumulation goes through per-batch scratch buffers reduced by a new `reduce_axis0` kernel (block tree-reduce, no atomicAdd). Same pattern as the existing LayerNorm bwd reducer — CUDA-Graph-safe. + +**Tech Stack:** Rust 1.85, CUDA 12.4, cudarc 0.19 (vendored), L40S target (`ci-training-l40s` pool). Stays on branch `ml-alpha-phase-a`. + +**Reference:** [Design spec](../specs/2026-05-17-kloop-parallelization-design.md) + +--- + +## File Map + +| File | Action | Responsibility | +|------|--------|----------------| +| `crates/ml-alpha/cuda/reduce_axis0.cu` | **CREATE** | Single parameterised reducer kernel `[B, N] → [N]`. Block tree-reduce across B. | +| `crates/ml-alpha/cuda/cfc_step.cu` | MODIFY | Block-per-batch refactor of `cfc_step_batched` fwd + bwd. Bwd writes to per-batch grad scratch. | +| `crates/ml-alpha/cuda/multi_horizon_heads.cu` | MODIFY | Block-per-batch refactor of `multi_horizon_heads_grn_bwd_batched`. 10 per-batch grad scratches. | +| `crates/ml-alpha/cuda/variable_selection.cu` | MODIFY | Block-per-batch refactor of `variable_selection_bwd`. 2 per-batch grad scratches. | +| `crates/ml-alpha/cuda/attention_pool.cu` | MODIFY | Block-per-batch refactor of `attention_pool_bwd`. 1 per-batch grad scratch (`grad_Q`). | +| `crates/ml-alpha/build.rs` | MODIFY | Register `reduce_axis0` kernel; bump cache-bust comment to v11. | +| `crates/ml-alpha/src/trainer/perception.rs` | MODIFY | 17 new scratch fields + `reduce_axis0_fn` + memset_zeros at step start + reducer launches after K-loop / after standalone bwd kernels. | +| `crates/ml-alpha/tests/cfc_bwd_b1_oracle.rs` | **CREATE** | New B=1 bit-equivalence test comparing `cfc_step_backward_batched` at B=1 vs `cfc_step_backward_gpu`. | +| `crates/ml-alpha/tests/perception_overfit.rs` | MODIFY | Add `stacked_trainer_loss_shrinks_at_batch_32` (B=32 smoke — first test that exercises cross-batch reduction). Add `cfc_bwd_scratch_clears_between_steps`. | + +**File responsibilities:** +- Each CUDA file owns one kernel-set. The block-per-batch refactor is local to each file. +- `reduce_axis0.cu` is the only new CUDA file; it's a generic reducer reused by all four refactored bwd kernels. +- `perception.rs` is the integration point — accumulates the wiring for all four refactors. +- Tests are isolated by concern: B=1 oracle, B=32 convergence smoke, scratch-clearing invariant. + +--- + +## Commit 1: `reduce_axis0` kernel + `cfc_step` refactor + +This is the load-bearing commit — establishes the scratch+reducer pattern and exercises it on the heaviest-hit K-loop kernel. Subsequent commits (GRN, VSN, attn) reuse the same `reduce_axis0` kernel + same trainer-wiring pattern. + +### Task 1: Write the `reduce_axis0` kernel + +**Files:** +- Create: `crates/ml-alpha/cuda/reduce_axis0.cu` +- Modify: `crates/ml-alpha/build.rs:10-22` (KERNELS list + cache-bust comment) + +- [ ] **Step 1: Create `crates/ml-alpha/cuda/reduce_axis0.cu`** + +```cuda +// reduce_axis0.cu — sum [B, N] → [N] along the leading axis. +// +// One block per output column (N blocks total). Threads tile B via a +// stride-loop, then block-wide tree-reduce. No atomicAdd per +// feedback_no_atomicadd.md. OVERWRITE semantics: out[j] = sum. +// +// Block dim is hard-coded to 256 — matches the typical tail dim +// granularity in the trainer (well within L40S limits, fits in shared +// memory at 1 KiB per block). + +#define REDAX0_BLOCK 256 + +extern "C" __global__ void reduce_axis0( + const float* __restrict__ per_batch, // [B, N] + int n_batch, + int n_tail, + float* __restrict__ out // [N] — OVERWRITE +) { + int j = blockIdx.x; + int tid = threadIdx.x; + if (j >= n_tail) return; + + __shared__ float s[REDAX0_BLOCK]; + + float my_sum = 0.0f; + // Stride-loop across B — handles n_batch > BLOCK and n_batch < BLOCK uniformly. + for (int bi = tid; bi < n_batch; bi += REDAX0_BLOCK) { + my_sum += per_batch[(long long)bi * n_tail + j]; + } + s[tid] = my_sum; + __syncthreads(); + + // Block tree-reduce. + for (int stride = REDAX0_BLOCK / 2; stride > 0; stride >>= 1) { + if (tid < stride) s[tid] += s[tid + stride]; + __syncthreads(); + } + + if (tid == 0) out[j] = s[0]; +} +``` + +- [ ] **Step 2: Register kernel in build.rs** + +In `crates/ml-alpha/build.rs`, modify the KERNELS slice (line 10-22). Append `"reduce_axis0"` after the last entry: + +```rust +const KERNELS: &[&str] = &[ + "mamba2_alpha_kernel", + "snap_feature_assemble", + "cfc_step", + "multi_horizon_heads", + "projection", + "bce_loss_multi_horizon", + "adamw_step", + "grad_norm", + "horizon_lambda", + "layer_norm", + "variable_selection", + "attention_pool", + "reduce_axis0", // Phase B: cross-batch param-grad reducer +]; +``` + +Bump the cache-bust comment to v11: + +```rust +// Cache bust v11 (2026-05-17): K-loop parallelization Phase B — +// new reduce_axis0.cu kernel + block-per-batch refactor of +// cfc_step_batched (fwd+bwd). Old cubins don't have the new symbols. +// Force fresh nvcc compile. +``` + +- [ ] **Step 3: Compile to verify the cubin builds** + +Run: +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo build -p ml-alpha --features cuda 2>&1 | tail -5 +``` + +Expected: `Finished \`dev\` profile [unoptimized + debuginfo] target(s)` with no errors. The new `reduce_axis0.cubin` artifact lives in OUT_DIR. + +### Task 2: Refactor `cfc_step_batched` forward kernel + +**Files:** +- Modify: `crates/ml-alpha/cuda/cfc_step.cu:185-213` + +- [ ] **Step 1: Replace the forward kernel body** + +In `crates/ml-alpha/cuda/cfc_step.cu`, replace the `cfc_step_batched` kernel (lines 185-213) with: + +```cuda +// Block-per-batch: grid=(n_batch, 1, 1), block=(n_hid, 1, 1). +// Each block handles ONE batch's n_hid channels in parallel. +// Removes the internal batch loop — frees ~31 SMs per launch at B=32. +extern "C" __global__ void cfc_step_batched( + const float* __restrict__ w_in, // [n_hid, n_in], row-major + const float* __restrict__ w_rec, // [n_hid, n_hid], row-major + const float* __restrict__ b, // [n_hid] + const float* __restrict__ tau, // [n_hid] + const float* __restrict__ x, // [n_batch, n_in] + const float* __restrict__ h_old, // [n_batch, n_hid] + float dt_s, + int n_in, + int n_hid, + int n_batch, + float* __restrict__ h_new // [n_batch, n_hid] +) { + int bi = blockIdx.x; + int i = threadIdx.x; + if (bi >= n_batch || i >= n_hid) return; + + const float bias_i = b[i]; + const float decay = expf(-dt_s / fmaxf(tau[i], 1e-6f)); + const float one_minus_decay = 1.0f - decay; + + const float* x_b = x + (long long)bi * n_in; + const float* h_old_b = h_old + (long long)bi * n_hid; + float pre = bias_i; + for (int k = 0; k < n_in; ++k) pre += w_in[i * n_in + k] * x_b[k]; + for (int k = 0; k < n_hid; ++k) pre += w_rec[i * n_hid + k] * h_old_b[k]; + h_new[(long long)bi * n_hid + i] = h_old_b[i] * decay + one_minus_decay * tanhf(pre); +} +``` + +### Task 3: Refactor `cfc_step_backward_batched` kernel signature + body + +**Files:** +- Modify: `crates/ml-alpha/cuda/cfc_step.cu:216-310+` (replace the entire `cfc_step_backward_batched` body) + +- [ ] **Step 1: Replace the backward kernel body** + +Replace the existing `cfc_step_backward_batched` (lines 216 onwards through its closing brace) with this block-per-batch version that writes to per-batch scratch: + +```cuda +// Block-per-batch backward. +// grid=(n_batch, 1, 1) block=(n_hid, 1, 1) +// shared mem: 2 * n_hid * 4 bytes (sd_pre + sdecay; one row each) +// +// Param-grad scratch tensors hold per-batch slices. Thread (bi, i) is +// sole writer to its slice of each scratch tensor for every K-iteration +// the K-loop calls this kernel — no race even with B parallel blocks. +// Caller zeroes the scratch buffers once per training step; this kernel +// uses += to accumulate across the K-loop's 64 invocations. +// +// Outputs grad_h_old and grad_x are per-batch indexed; each block bi is +// sole writer to its [bi, :] slice (overwrite, single launch). +extern "C" __global__ void cfc_step_backward_batched( + const float* __restrict__ w_in, // [n_hid, n_in] + const float* __restrict__ w_rec, // [n_hid, n_hid] + const float* __restrict__ b, // [n_hid] + const float* __restrict__ tau, // [n_hid] + const float* __restrict__ x, // [n_batch, n_in] + const float* __restrict__ h_old, // [n_batch, n_hid] + const float* __restrict__ grad_h_new, // [n_batch, n_hid] + float dt_s, + int n_in, + int n_hid, + int n_batch, + float* __restrict__ grad_w_in_scratch, // [n_batch, n_hid, n_in] accum += + float* __restrict__ grad_w_rec_scratch, // [n_batch, n_hid, n_hid] accum += + float* __restrict__ grad_b_scratch, // [n_batch, n_hid] accum += + float* __restrict__ grad_tau_scratch, // [n_batch, n_hid] accum += + float* __restrict__ grad_h_old, // [n_batch, n_hid] overwrite + float* __restrict__ grad_x // [n_batch, n_in] overwrite (nullptr OK) +) { + extern __shared__ float smem[]; + float* sd_pre = smem; // [n_hid] — one row for this block's bi + float* sdecay = sd_pre + n_hid; // [n_hid] + + int bi = blockIdx.x; + int i = threadIdx.x; + if (bi >= n_batch || i >= n_hid) return; + + const float bias_i = b[i]; + const float tau_eps = 1e-6f; + const float tau_eff = fmaxf(tau[i], tau_eps); + const float decay = expf(-dt_s / tau_eff); + sdecay[i] = decay; + + const float* x_b = x + (long long)bi * n_in; + const float* h_old_b = h_old + (long long)bi * n_hid; + const float dh_b = grad_h_new[(long long)bi * n_hid + i]; + + float pre = bias_i; + for (int k = 0; k < n_in; ++k) pre += w_in[i * n_in + k] * x_b[k]; + for (int k = 0; k < n_hid; ++k) pre += w_rec[i * n_hid + k] * h_old_b[k]; + const float s = tanhf(pre); + const float d_pre_b = dh_b * (1.0f - decay) * (1.0f - s * s); + sd_pre[i] = d_pre_b; + + // Pass 1: per-batch param-grad scratch writes. + // Single-writer per (bi, i, *) — block bi is sole writer to its slice; + // K-loop's 64 invocations accumulate within the same slice via +=. + grad_b_scratch[(long long)bi * n_hid + i] += d_pre_b; + for (int k = 0; k < n_in; ++k) { + grad_w_in_scratch[((long long)bi * n_hid + i) * n_in + k] += d_pre_b * x_b[k]; + } + for (int k = 0; k < n_hid; ++k) { + grad_w_rec_scratch[((long long)bi * n_hid + i) * n_hid + k] += d_pre_b * h_old_b[k]; + } + if (tau[i] > tau_eps) { + const float d_decay_b = dh_b * (h_old_b[i] - s); + grad_tau_scratch[(long long)bi * n_hid + i] += d_decay_b * decay * dt_s / (tau_eff * tau_eff); + } + __syncthreads(); + + // Pass 2: grad_h_old[bi, i] = sum_j sd_pre[j] * W_rec[j, i] + dh_b * decay. + { + float gh = dh_b * decay; + for (int j = 0; j < n_hid; ++j) { + gh += sd_pre[j] * w_rec[j * n_hid + i]; + } + grad_h_old[(long long)bi * n_hid + i] = gh; + } + + // Pass 3: grad_x[bi, k] via thread 0 of this block. + if (grad_x != nullptr && i == 0) { + float* grad_x_b = grad_x + (long long)bi * n_in; + for (int k = 0; k < n_in; ++k) { + float gx = 0.0f; + for (int j = 0; j < n_hid; ++j) { + gx += sd_pre[j] * w_in[j * n_in + k]; + } + grad_x_b[k] = gx; + } + } +} +``` + +- [ ] **Step 2: Compile to verify nvcc accepts the rewrite** + +Run: +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo build -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. The Rust side will fail because the trainer still calls the old kernel signature — that's wired in the next task. The cubin compiles independently. + +### Task 4: Add cfc scratch buffers + reducer fn to `PerceptionTrainer` + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/perception.rs` (struct fields + new() + dispatch_train_step) + +- [ ] **Step 1: Add cubin import** + +Near the top of `crates/ml-alpha/src/trainer/perception.rs`, find the existing `ATTENTION_POOL_CUBIN` line and add directly after it: + +```rust +const REDUCE_AXIS0_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/reduce_axis0.cubin")); +``` + +- [ ] **Step 2: Add new fields to `PerceptionTrainer` struct** + +Find the struct definition (search for `pub struct PerceptionTrainer {`). At the end of the struct (before the closing `}`), add a new section: + +```rust + // ── K-loop parallelization (Phase B) ── + // Per-batch grad scratch buffers for cfc_step_backward_batched. + // Zeroed once per training step; the K-loop's 64 bwd calls accumulate + // into these via +=. After the K-loop, reduce_axis0 sums across B + // into the final grad buffers (overwrite semantics). + cfc_grad_w_in_scratch_d: CudaSlice, // [B, n_hid, n_in] + cfc_grad_w_rec_scratch_d: CudaSlice, // [B, n_hid, n_hid] + cfc_grad_b_scratch_d: CudaSlice, // [B, n_hid] + cfc_grad_tau_scratch_d: CudaSlice, // [B, n_hid] + /// Cross-batch reducer kernel: `[B, N] → [N]` via block tree-reduce. + /// Used for every per-batch grad scratch in the refactored bwd path. + reduce_axis0_fn: CudaFunction, + /// Owns the reduce_axis0 cubin lifetime. + _reduce_module: Arc, +``` + +- [ ] **Step 3: Load the reduce_axis0 cubin in `PerceptionTrainer::new`** + +Find the existing cubin loads (search for `let ln_module = ctx`). After the last cubin load (likely `attn_module`), add: + +```rust + let reduce_module = ctx + .load_cubin(REDUCE_AXIS0_CUBIN.to_vec()) + .context("reduce_axis0 cubin")?; + let reduce_axis0_fn = reduce_module + .load_function("reduce_axis0") + .context("reduce_axis0 symbol")?; +``` + +- [ ] **Step 4: Allocate cfc scratch buffers in `PerceptionTrainer::new`** + +Find the existing CfC weight init block (search for `// CfC weights`). After the existing `let opt_tau = AdamW::new(...)` line in that block, add: + +```rust + let cfc_grad_w_in_scratch_d = stream + .alloc_zeros::(cfg.n_batch * n_hid * n_in) + .context("cfc_grad_w_in_scratch_d alloc")?; + let cfc_grad_w_rec_scratch_d = stream + .alloc_zeros::(cfg.n_batch * n_hid * n_hid) + .context("cfc_grad_w_rec_scratch_d alloc")?; + let cfc_grad_b_scratch_d = stream + .alloc_zeros::(cfg.n_batch * n_hid) + .context("cfc_grad_b_scratch_d alloc")?; + let cfc_grad_tau_scratch_d = stream + .alloc_zeros::(cfg.n_batch * n_hid) + .context("cfc_grad_tau_scratch_d alloc")?; +``` + +- [ ] **Step 5: Populate the new fields in `Ok(Self { ... })`** + +Find the `Ok(Self {` block at the end of `new()`. Add these fields anywhere in the struct (matching the field declaration order): + +```rust + cfc_grad_w_in_scratch_d, + cfc_grad_w_rec_scratch_d, + cfc_grad_b_scratch_d, + cfc_grad_tau_scratch_d, + reduce_axis0_fn, + _reduce_module: reduce_module, +``` + +- [ ] **Step 6: Build to verify the trainer compiles with new fields** + +Run: +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo check -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. Fields are declared but unused (dead-code warnings on `reduce_axis0_fn` etc. are OK at this checkpoint — wiring comes next). + +### Task 5: Update `dispatch_train_step` — cfc fwd launch + scratch memset + bwd launch + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/perception.rs` (inside `dispatch_train_step`) + +- [ ] **Step 1: Change cfc forward launch config** + +In `dispatch_train_step`, find `let cfg_cfc = LaunchConfig {` (around line 1130-ish — the launch config for the K-loop CfC step). Replace its definition with: + +```rust + // CfC fwd: block-per-batch (Phase B). One block per sample; + // each block's threads parallelize over n_hid channels. + let cfg_cfc = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (HIDDEN_DIM as u32, 1, 1), + shared_mem_bytes: 0, + }; +``` + +- [ ] **Step 2: Change cfc backward launch config** + +Find `let cfg_cfc_bwd = LaunchConfig {` (a few lines after `cfg_cfc`). Replace with: + +```rust + // CfC bwd: block-per-batch + 2 * n_hid floats of shared mem + // (one sd_pre row + sdecay for the block's single bi). + let cfc_bwd_smem = (2 * 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, + }; +``` + +- [ ] **Step 3: Replace cfc-grad memset_zeros with scratch memset_zeros** + +Find the block `// Zero shared param-grad accumulators (kernels += into these).` Inside that block, replace the four cfc grad zeroings (lines starting with `self.stream.memset_zeros(&mut self.grad_w_in_d)` etc.) with: + +```rust + // CfC per-batch grad scratch: zero ONCE per step; K-loop bwd + // accumulates into these, then reduce_axis0 collapses → final + // grad buffers (OVERWRITE) after the K-loop. AdamW reads the + // final grads (must run AFTER the reducer). + self.stream.memset_zeros(&mut self.cfc_grad_w_in_scratch_d) + .map_err(|e| anyhow::anyhow!("zero cfc_grad_w_in_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.cfc_grad_w_rec_scratch_d) + .map_err(|e| anyhow::anyhow!("zero cfc_grad_w_rec_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.cfc_grad_b_scratch_d) + .map_err(|e| anyhow::anyhow!("zero cfc_grad_b_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.cfc_grad_tau_scratch_d) + .map_err(|e| anyhow::anyhow!("zero cfc_grad_tau_scratch: {e}"))?; +``` + +DELETE the previous lines that zeroed `grad_w_in_d`, `grad_w_rec_d`, `grad_b_d`, `grad_tau_d`. Keep the other zeroings (GRN, LN, etc.) — they are untouched by this commit. + +- [ ] **Step 4: Update cfc backward kernel call to pass scratch buffers** + +Find the cfc backward launch inside the K-loop (search for `step_bwd_batched_fn`). Replace the launch block (the one that sets up args via `.arg(...)`) with: + +```rust + // cfc_step_bwd_batched: per-batch param-grad scratch + per-batch + // grad_h_old/grad_x outputs. Writes to scratch via += across K + // iterations; reduce_axis0 collapses B → final grad after K-loop. + unsafe { + let mut launch = self.stream.launch_builder(&self.step_bwd_batched_fn); + launch + .arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_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")?; + } +``` + +The kernel signature changed: the old `grad_w_in_d, grad_w_rec_d, grad_b_d, grad_tau_d` direct args are replaced with the four `*_scratch_d` buffers. + +- [ ] **Step 5: Add reducer launches after the K-loop (before AdamW)** + +Find the end of the bwd K-loop (search for `} // end of K-loop bwd` or look for where the loop body's closing `}` is, then the section right after — typically the AdamW step block starts with `// ── 9. Apply AdamW updates`). + +Just BEFORE the AdamW block, add: + +```rust + // ── 8c. Reduce cfc per-batch grad scratch → final grad buffers. + // Block tree-reduce across bi; deterministic sum order. + // Each launch is grid=(n_tail, 1, 1), block=(256, 1, 1). + // Total: 4 launches, each handling one cfc param tensor. + { + let n_batch_i = b_sz as i32; + let reduce_block: u32 = 256; + let reduce_at = |n_tail: usize, + scratch: &CudaSlice, + out: &mut CudaSlice, + label: &'static str| + -> Result<()> { + let cfg = LaunchConfig { + grid_dim: (n_tail as u32, 1, 1), + block_dim: (reduce_block, 1, 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_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")?; + } +``` + +Note: `cfc_step_batched`'s w_in is `[n_hid, n_in]` where `n_in = HIDDEN_DIM` (CfC input is the LN_b output, also HIDDEN_DIM). So `n_tail = n_hid * n_in = HIDDEN_DIM * HIDDEN_DIM = 16384`. + +- [ ] **Step 6: Build to verify all changes compile** + +Run: +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo check -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build, no errors. Some dead-code warnings on `grad_w_in_d`/etc. are now stale (these ARE used — they're the reducer outputs) but rust-analyzer's warnings can lag; cargo check is authoritative. + +### Task 6: B=1 oracle test for cfc backward + +**Files:** +- Create: `crates/ml-alpha/tests/cfc_bwd_b1_oracle.rs` + +- [ ] **Step 1: Write the test** + +Create `crates/ml-alpha/tests/cfc_bwd_b1_oracle.rs`: + +```rust +//! Cfc bwd B=1 bit-equivalence-tolerant oracle. +//! +//! Compares the refactored block-per-batch `cfc_step_backward_batched` +//! kernel at B=1 against the established single-sample helper +//! `cfc_step_backward_gpu`. Asserts agreement within FP tolerance +//! (relative_eq ≤ 1e-6 — NOT bit-exact; the two kernels may sum +//! in different orders, and FP addition is non-associative). +//! +//! This test exists ONLY for cfc_step. GRN/VSN/attn bwd kernels have no +//! single-sample GPU oracle (per spec section "B=1 oracle asymmetry"); +//! they rely on convergence smoke + cluster A/B AUC trajectory match. + +use approx::assert_relative_eq; +use cudarc::driver::{DevicePtr, DevicePtrMut, LaunchConfig, PushKernelArg}; +use ml_alpha::cfc::step::{cfc_step_backward_gpu, CfcWeights}; +use ml_core::device::MlDevice; +use rand::{Rng, SeedableRng}; +use rand_chacha::ChaCha8Rng; +use std::sync::Arc; + +const CFC_STEP_CUBIN: &[u8] = include_bytes!( + concat!(env!("OUT_DIR"), "/cfc_step.cubin") +); +const REDUCE_AXIS0_CUBIN: &[u8] = include_bytes!( + concat!(env!("OUT_DIR"), "/reduce_axis0.cubin") +); + +fn test_device() -> MlDevice { + MlDevice::cuda(0).expect("CUDA 0 required for ml-alpha tests") +} + +fn rand_weights(seed: u64, n_in: usize, n_hid: usize) -> CfcWeights { + let mut r = ChaCha8Rng::seed_from_u64(seed); + CfcWeights { + w_in: (0..n_hid * n_in).map(|_| r.gen_range(-0.1..0.1)).collect(), + w_rec: (0..n_hid * n_hid).map(|_| r.gen_range(-0.05..0.05)).collect(), + b: (0..n_hid).map(|_| r.gen_range(-0.01..0.01)).collect(), + tau: (0..n_hid).map(|_| r.gen_range(0.05..2.0)).collect(), + n_in, + n_hid, + } +} + +#[test] +fn cfc_bwd_b1_matches_single_sample() { + let dev = test_device(); + let stream = dev.cuda_stream().expect("stream").clone(); + let ctx = dev.cuda_context().expect("ctx"); + + let n_in = 8usize; + let n_hid = 16usize; + let dt_s = 0.02_f32; + let w = rand_weights(0xABCDE, n_in, n_hid); + + let mut r = ChaCha8Rng::seed_from_u64(0xFADE); + let x: Vec = (0..n_in).map(|_| r.gen_range(-1.0..1.0)).collect(); + let h_old: Vec = (0..n_hid).map(|_| r.gen_range(-1.0..1.0)).collect(); + let grad_h_new: Vec = (0..n_hid).map(|_| r.gen_range(-0.5..0.5)).collect(); + + // Oracle: single-sample helper. + let (g_w_in_ref, g_w_rec_ref, g_b_ref, g_tau_ref, _g_h_old_ref) = + cfc_step_backward_gpu(&dev, &w, &x, &h_old, &grad_h_new, dt_s).expect("oracle"); + + // Subject: launch the refactored batched kernel at B=1, run reducer, + // read back final grads. The reducer at B=1 is a copy — sum over a + // single batch slice equals that slice. + let cfc_module = ctx.load_cubin(CFC_STEP_CUBIN.to_vec()).expect("cfc cubin"); + let bwd_fn = cfc_module.load_function("cfc_step_backward_batched").expect("bwd symbol"); + + let reduce_module = ctx.load_cubin(REDUCE_AXIS0_CUBIN.to_vec()).expect("reduce cubin"); + let reduce_fn = reduce_module.load_function("reduce_axis0").expect("reduce symbol"); + + let upload = |slice: &[f32]| -> cudarc::driver::CudaSlice { + let h = cudarc::driver::result::malloc_host::(slice.len()).unwrap(); + unsafe { std::slice::from_raw_parts_mut(h as *mut f32, slice.len()) } + .copy_from_slice(slice); + // Easier: use stream.memcpy_stod via Rust slice + stream.memcpy_stod(slice).expect("upload") + }; + + let w_in_d = upload(&w.w_in); + let w_rec_d = upload(&w.w_rec); + let b_d = upload(&w.b); + let tau_d = upload(&w.tau); + let x_d = upload(&x); + let h_old_d = upload(&h_old); + let grad_h_new_d = upload(&grad_h_new); + + let n_batch = 1usize; + let mut grad_w_in_scratch_d = stream.alloc_zeros::(n_batch * n_hid * n_in).expect("scratch"); + let mut grad_w_rec_scratch_d = stream.alloc_zeros::(n_batch * n_hid * n_hid).expect("scratch"); + let mut grad_b_scratch_d = stream.alloc_zeros::(n_batch * n_hid).expect("scratch"); + let mut grad_tau_scratch_d = stream.alloc_zeros::(n_batch * n_hid).expect("scratch"); + let mut grad_h_old_d = stream.alloc_zeros::(n_batch * n_hid).expect("grad_h_old"); + let mut grad_x_d = stream.alloc_zeros::(n_batch * n_in).expect("grad_x"); + + let smem = (2 * n_hid * std::mem::size_of::()) as u32; + let cfg = LaunchConfig { + grid_dim: (n_batch as u32, 1, 1), + block_dim: (n_hid as u32, 1, 1), + shared_mem_bytes: smem, + }; + let n_in_i = n_in as i32; + let n_hid_i = n_hid as i32; + let n_batch_i = n_batch as i32; + unsafe { + let mut launch = stream.launch_builder(&bwd_fn); + launch + .arg(&w_in_d).arg(&w_rec_d).arg(&b_d).arg(&tau_d) + .arg(&x_d).arg(&h_old_d).arg(&grad_h_new_d) + .arg(&dt_s).arg(&n_in_i).arg(&n_hid_i).arg(&n_batch_i) + .arg(&mut grad_w_in_scratch_d).arg(&mut grad_w_rec_scratch_d) + .arg(&mut grad_b_scratch_d).arg(&mut grad_tau_scratch_d) + .arg(&mut grad_h_old_d).arg(&mut grad_x_d); + launch.launch(cfg).expect("bwd launch"); + } + + // Reduce: at B=1 the reducer is a copy. + let mut grad_w_in_d = stream.alloc_zeros::(n_hid * n_in).expect("final"); + let mut grad_w_rec_d = stream.alloc_zeros::(n_hid * n_hid).expect("final"); + let mut grad_b_d = stream.alloc_zeros::(n_hid).expect("final"); + let mut grad_tau_d = stream.alloc_zeros::(n_hid).expect("final"); + + let reduce = |scratch: &cudarc::driver::CudaSlice, + out: &mut cudarc::driver::CudaSlice, + n_tail: usize| { + let cfg = LaunchConfig { + grid_dim: (n_tail as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }; + let n_tail_i = n_tail as i32; + unsafe { + let mut launch = stream.launch_builder(&reduce_fn); + launch.arg(scratch).arg(&n_batch_i).arg(&n_tail_i).arg(out); + launch.launch(cfg).expect("reduce launch"); + } + }; + reduce(&grad_w_in_scratch_d, &mut grad_w_in_d, n_hid * n_in); + reduce(&grad_w_rec_scratch_d, &mut grad_w_rec_d, n_hid * n_hid); + reduce(&grad_b_scratch_d, &mut grad_b_d, n_hid); + reduce(&grad_tau_scratch_d, &mut grad_tau_d, n_hid); + + stream.synchronize().expect("sync"); + + let g_w_in = stream.memcpy_dtov(&grad_w_in_d).expect("dl"); + let g_w_rec = stream.memcpy_dtov(&grad_w_rec_d).expect("dl"); + let g_b = stream.memcpy_dtov(&grad_b_d).expect("dl"); + let g_tau = stream.memcpy_dtov(&grad_tau_d).expect("dl"); + + // Compare element-wise to oracle within FP tolerance. + assert_eq!(g_w_in.len(), g_w_in_ref.len()); + for (i, (&a, &b)) in g_w_in.iter().zip(g_w_in_ref.iter()).enumerate() { + assert_relative_eq!(a, b, max_relative = 1e-5, epsilon = 1e-6, + "grad_w_in[{i}] mismatch"); + } + for (i, (&a, &b)) in g_w_rec.iter().zip(g_w_rec_ref.iter()).enumerate() { + assert_relative_eq!(a, b, max_relative = 1e-5, epsilon = 1e-6, + "grad_w_rec[{i}] mismatch"); + } + for (i, (&a, &b)) in g_b.iter().zip(g_b_ref.iter()).enumerate() { + assert_relative_eq!(a, b, max_relative = 1e-5, epsilon = 1e-6, + "grad_b[{i}] mismatch"); + } + for (i, (&a, &b)) in g_tau.iter().zip(g_tau_ref.iter()).enumerate() { + assert_relative_eq!(a, b, max_relative = 1e-5, epsilon = 1e-6, + "grad_tau[{i}] mismatch"); + } +} +``` + +Note: the helper-import name `upload` is local to the test; uses `stream.memcpy_stod` (cudarc 0.19 API). If that exact API name doesn't exist, fall back to allocating a `CudaSlice` and using `stream.memcpy_htod`. The actual upload pattern is standard cudarc; check `crates/ml-alpha/src/cfc/snap_features.rs::snap_feature_assemble_gpu` for the canonical example. + +- [ ] **Step 2: Run the test** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test cfc_bwd_b1_oracle -- --nocapture 2>&1 | tail -20 +``` + +Expected: `test cfc_bwd_b1_matches_single_sample ... ok`. + +### Task 7: Add B=32 smoke test in perception_overfit + +**Files:** +- Modify: `crates/ml-alpha/tests/perception_overfit.rs` (append new test) + +- [ ] **Step 1: Add the test** + +Find the existing test `stacked_trainer_loss_shrinks_with_stride_4` and append the new test immediately after: + +```rust +/// Verifies the trainer converges at n_batch=32 — the FIRST test that +/// exercises the cross-batch reducer code path. Existing +/// `stacked_trainer_loss_shrinks_*` tests all use n_batch=1 so the new +/// per-batch scratch + reducer logic was never previously hit. +#[test] +fn stacked_trainer_loss_shrinks_at_batch_32() { + let dev = test_device(); + let cfg = PerceptionTrainerConfig { + seq_len: 16, + mamba2_state_dim: 8, + lr_cfc: 3e-3, + lr_mamba2: 1e-3, + seed: 0xB32B, + horizon_weights: [1.0; 5], + n_batch: 32, + decision_stride: 1, + }; + let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init"); + + let mut initial = 0.0_f32; + let mut ts_base = 1_000_000u64; + let mut prev_mid = 5500.0_f32; + let mut warm_seqs: Vec> = Vec::with_capacity(cfg.n_batch); + let mut warm_labels: Vec> = Vec::with_capacity(cfg.n_batch); + for warmup in 0..4 { + warm_seqs.clear(); + warm_labels.clear(); + for _ in 0..cfg.n_batch { + let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts_base); + prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]); + ts_base = seq.last().unwrap().ts_ns; + warm_seqs.push(seq); + warm_labels.push(labels); + } + let seq_refs: Vec<&[Mbp10RawInput]> = warm_seqs.iter().map(|s| s.as_slice()).collect(); + let lbl_refs: Vec<&[[f32; 5]]> = warm_labels.iter().map(|l| l.as_slice()).collect(); + let l = trainer.step_batched(&seq_refs, &lbl_refs).expect("warm step"); + if warmup >= 2 { + initial += l; + } + } + initial /= 2.0; + eprintln!("B=32 trainer: initial={initial:.4}"); + + // Train 200 steps. + for _ in 0..200 { + let mut seqs: Vec> = Vec::with_capacity(cfg.n_batch); + let mut labels: Vec> = Vec::with_capacity(cfg.n_batch); + for _ in 0..cfg.n_batch { + let (seq, lbl) = synthetic_seq(cfg.seq_len, prev_mid, ts_base); + prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]); + ts_base = seq.last().unwrap().ts_ns; + seqs.push(seq); + labels.push(lbl); + } + let seq_refs: Vec<&[Mbp10RawInput]> = seqs.iter().map(|s| s.as_slice()).collect(); + let lbl_refs: Vec<&[[f32; 5]]> = labels.iter().map(|l| l.as_slice()).collect(); + trainer.step_batched(&seq_refs, &lbl_refs).expect("train step"); + } + + let mut final_loss = 0.0_f32; + for _ in 0..4 { + let mut seqs: Vec> = Vec::with_capacity(cfg.n_batch); + let mut labels: Vec> = Vec::with_capacity(cfg.n_batch); + for _ in 0..cfg.n_batch { + let (seq, lbl) = synthetic_seq(cfg.seq_len, prev_mid, ts_base); + prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]); + ts_base = seq.last().unwrap().ts_ns; + seqs.push(seq); + labels.push(lbl); + } + let seq_refs: Vec<&[Mbp10RawInput]> = seqs.iter().map(|s| s.as_slice()).collect(); + let lbl_refs: Vec<&[[f32; 5]]> = labels.iter().map(|l| l.as_slice()).collect(); + final_loss += trainer.step_batched(&seq_refs, &lbl_refs).expect("final step"); + } + final_loss /= 4.0; + eprintln!("B=32 trainer: final={final_loss:.4}"); + assert!( + final_loss < 0.6 * initial || final_loss < 0.1, + "B=32 trainer failed to converge: {initial:.4} → {final_loss:.4}" + ); +} +``` + +- [ ] **Step 2: Run the B=32 smoke** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit stacked_trainer_loss_shrinks_at_batch_32 -- --nocapture 2>&1 | tail -15 +``` + +Expected: `test stacked_trainer_loss_shrinks_at_batch_32 ... ok`, final loss < 0.1. + +### Task 8: Add scratch-clears test + +**Files:** +- Modify: `crates/ml-alpha/tests/perception_overfit.rs` (append) + +- [ ] **Step 1: Add a public accessor for scratch readback (test-only)** + +In `crates/ml-alpha/src/trainer/perception.rs`, find the existing test-diagnostic accessors (search for `pub fn loss_ema_snapshot`). Add after them: + +```rust + /// Test/diagnostic readback of the cfc per-batch scratch buffer for + /// grad_b. Used by the scratch-clears test to verify the scratch is + /// zeroed at the start of each training step. Forces a stream sync. + pub fn cfc_grad_b_scratch_snapshot(&self) -> Result> { + let n = self.cfg.n_batch * HIDDEN_DIM; + let staging = unsafe { crate::pinned_mem::MappedF32Buffer::new(n) } + .map_err(|e| anyhow::anyhow!("scratch staging: {e}"))?; + unsafe { + let (src_ptr, _g) = self.cfc_grad_b_scratch_d.device_ptr(&self.stream); + cudarc::driver::result::memcpy_dtod_async( + staging.dev_ptr, src_ptr, + n * std::mem::size_of::(), + self.stream.cu_stream(), + ).context("scratch dtod")?; + } + self.stream.synchronize().context("scratch sync")?; + Ok(staging.read_all()) + } +``` + +- [ ] **Step 2: Add the test in perception_overfit.rs** + +Append after the B=32 test: + +```rust +/// Asserts cfc per-batch scratch is zero at the start of step N+1 — +/// i.e., it doesn't retain step N's accumulated gradient. Guards +/// against the "forgot to zero scratch" class of regression. +#[test] +fn cfc_bwd_scratch_clears_between_steps() { + let dev = test_device(); + let cfg = PerceptionTrainerConfig { + seq_len: 16, + mamba2_state_dim: 8, + lr_cfc: 3e-3, + lr_mamba2: 1e-3, + seed: 0x5C0, + horizon_weights: [1.0; 5], + n_batch: 1, + decision_stride: 1, + }; + let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init"); + + // Step 1: run a normal training step. Scratch will have non-zero + // contributions from K-loop bwd. + let (seq_a, labels_a) = synthetic_seq(cfg.seq_len, 5500.0, 1_000_000); + let _ = trainer.step(&seq_a, &labels_a).expect("step A"); + + // After AdamW + reducer, the scratch buffer is what the trainer + // captured during step A's K-loop. It's NOT zeroed by the AdamW path. + // The zero happens at the START of the next step. + + // Step 2: run another step. Inside dispatch_train_step, scratch is + // memset_zeros'd BEFORE the K-loop bwd begins. If we could pause + // mid-step we'd see the scratch as zero. Instead, the canonical + // way to assert this is by training-equivalence: a step whose + // grad contributions came from a different sequence should NOT + // be polluted by step A's residual gradient. + // + // The cleanest assertion: run step B with synthetic_seq starting + // at a DIFFERENT prev_mid (so labels and gradients differ from + // step A); the trainer must produce a non-degenerate loss for step + // B (not zero, not nan). If scratch retained step A's contribution, + // step B's first-K-iteration scratch read would observe a non-zero + // start state, which would manifest as NaN propagation given that + // the post-reducer grad_*_d is overwritten — but the per-batch + // scratch is what the kernel reads via += semantics. + // + // Concrete check: after step A, snapshot the scratch. Train step B. + // Snapshot again. The post-step-B scratch holds ONLY step B's + // contribution (memset_zeros at start of step B wiped step A out). + // We assert step-B scratch != step-A scratch (different gradients). + let snap_a = trainer.cfc_grad_b_scratch_snapshot().expect("snap A"); + let (seq_b, labels_b) = synthetic_seq(cfg.seq_len, 5800.0, 2_000_000); + let _ = trainer.step(&seq_b, &labels_b).expect("step B"); + let snap_b = trainer.cfc_grad_b_scratch_snapshot().expect("snap B"); + + // The two scratch snapshots come from different training sequences; + // they MUST differ in at least one element. If they're equal, + // either scratch wasn't zeroed (it retained step A) or something + // else is wrong (e.g., the K-loop didn't run, or the bwd kernel + // is a no-op). + let differs = snap_a.iter().zip(snap_b.iter()).any(|(a, b)| (a - b).abs() > 1e-8); + assert!( + differs, + "scratch did not change between two different training steps — \ + this suggests scratch is not being zeroed at step start" + ); +} +``` + +- [ ] **Step 3: Run the scratch-clears test** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit cfc_bwd_scratch_clears_between_steps -- --nocapture 2>&1 | tail -10 +``` + +Expected: `test cfc_bwd_scratch_clears_between_steps ... ok`. + +### Task 9: Run all existing perception_overfit tests + +- [ ] **Step 1: Run the full perception_overfit suite** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit -- --nocapture 2>&1 | tail -30 +``` + +Expected: all tests pass (8 existing + 2 new = 10). Critical: `stacked_trainer_loss_shrinks_on_constant_signal` and `stacked_trainer_loss_shrinks_with_stride_4` both still converge to ~0 in 250 steps. + +If they regress, the refactor has a chain-rule bug. The B=1 oracle test (Task 6) should have caught it; if not, the bug is in code paths the oracle doesn't exercise (e.g., the K-loop scratch accumulation across iterations). + +- [ ] **Step 2: Run finite-diff tests** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test backward_finite_diff -- --nocapture 2>&1 | tail -15 +``` + +Expected: all finite-diff tests pass. These use the single-sample GPU helpers (orthogonal code path) — they're unchanged by this refactor but verifying confirms the kernel module loading is intact. + +### Task 10: Commit Commit 1 + +- [ ] **Step 1: Stage all files** + +```bash +git add crates/ml-alpha/cuda/reduce_axis0.cu \ + crates/ml-alpha/cuda/cfc_step.cu \ + crates/ml-alpha/build.rs \ + crates/ml-alpha/src/trainer/perception.rs \ + crates/ml-alpha/tests/cfc_bwd_b1_oracle.rs \ + crates/ml-alpha/tests/perception_overfit.rs +``` + +- [ ] **Step 2: Commit** + +```bash +git commit -m "$(cat <<'EOF' +perf(ml-alpha): block-per-batch cfc_step + reduce_axis0 reducer (K-loop Phase B commit 1) + +Per docs/superpowers/specs/2026-05-17-kloop-parallelization-design.md. + +cfc_step_batched (fwd + bwd) refactored from grid=(1,1,1) with internal +n_batch loop to grid=(B,1,1) — each block handles one batch. Removes +the single-SM bottleneck on the K-loop's most-called kernel (64×/step). + +Param-grad accumulation moves to per-batch scratch: + cfc_grad_w_in_scratch_d [B, n_hid, n_in] + cfc_grad_w_rec_scratch_d [B, n_hid, n_hid] + cfc_grad_b_scratch_d [B, n_hid] + cfc_grad_tau_scratch_d [B, n_hid] + +Zeroed once per training step, K-loop's 64 bwd calls += into them, then +4 reduce_axis0 launches collapse B → final grad buffers (OVERWRITE) before +AdamW. New AdamW-after-reducer invariant: final grads are meaningful only +after the reducer has run in the current step. + +New reduce_axis0 kernel: single parameterised reducer [B, N] → [N] via +block tree-reduce (no atomicAdd per feedback_no_atomicadd.md). Same +pattern as layer_norm_reduce_param_grads — CUDA-Graph-safe. + +Tests: +- New cfc_bwd_b1_oracle.rs: at B=1 the refactored bwd matches the + single-sample helper cfc_step_backward_gpu within relative_eq 1e-5 + (FP-tolerant, not bit-exact — different kernels may sum in different + orders). +- New stacked_trainer_loss_shrinks_at_batch_32: FIRST test that + actually exercises the cross-batch reduction code path; existing + perception_overfit suite was all B=1. +- New cfc_bwd_scratch_clears_between_steps: guards against the "forgot + to zero scratch" regression class. + +build.rs: +- KERNELS list adds "reduce_axis0" +- Cache-bust → v11 + +Co-Authored-By: Claude Opus 4.7 +EOF +)" +``` + +Expected output: `[ml-alpha-phase-a ] perf(ml-alpha): block-per-batch ...` with no pre-commit failures. + +- [ ] **Step 3: Push** + +```bash +git push origin ml-alpha-phase-a 2>&1 | tail -3 +``` + +Expected: push succeeds, no force-push needed. + +--- + +## Commit 2: GRN bwd refactor + +Same pattern as Commit 1, applied to `multi_horizon_heads_grn_bwd_batched`. Ten parameter tensors → ten scratch buffers → ten reducer launches. + +### Task 11: Refactor GRN bwd kernel signature + body + +**Files:** +- Modify: `crates/ml-alpha/cuda/multi_horizon_heads.cu` (the `multi_horizon_heads_grn_bwd_batched` function) + +- [ ] **Step 1: Locate the existing kernel** + +The kernel was added in commit `5e23005de` ("TFT GRN forward+backward kernels for multi-horizon heads"). Find `extern "C" __global__ void multi_horizon_heads_grn_bwd_batched(` in `crates/ml-alpha/cuda/multi_horizon_heads.cu`. + +- [ ] **Step 2: Replace the kernel body** + +Replace the entire `multi_horizon_heads_grn_bwd_batched` function with the block-per-batch version. Key changes from the existing kernel: + +1. Launch contract changes from `grid=(1,1,1)` to `grid=(n_batch,1,1)`. +2. Outer `for (int bi = 0; bi < n_batch; ++bi)` loop removed; `int bi = blockIdx.x;` replaces it. +3. Every param-grad `+=` now targets a per-batch scratch slice. Indexing for `grad_w1[k*HEAD_MID*HIDDEN + m*HIDDEN + i] += ...` becomes: + ```cuda + grad_w1_scratch[((bi * N_HORIZONS_H + k) * HEAD_MID_H + m) * HIDDEN_H + i] += ... + ``` + Same pattern for `grad_w2`, `grad_b1`, `grad_b2`, `grad_w_gate`, `grad_b_gate`, `grad_w_main`, `grad_b_main`, `grad_w_skip`, `grad_b_skip`. +4. The `s_lambda` shared array is loaded by threads 0..N_HORIZONS_H within EACH block (was once for the whole launch). With one block per batch, the load is cheap and the broadcast pattern stays the same. +5. Shared-mem usage per block drops because there's no batch loop — but the existing kernel already had no per-bi shared accumulator, so this is unchanged. + +The full kernel text replaces the existing body. Mechanical translation rules above. Reference the existing kernel for the chain-rule math; only the indexing and launch contract change. + +Signature change for the kernel (rename of the param-grad output args makes the new role explicit): + +```cuda +extern "C" __global__ void multi_horizon_heads_grn_bwd_batched( + const float* __restrict__ w1, + const float* __restrict__ w2, + const float* __restrict__ w_gate, + const float* __restrict__ w_main, + const float* __restrict__ w_skip, + const float* __restrict__ probs, + const float* __restrict__ grad_probs, + const float* __restrict__ z1, + const float* __restrict__ a1, + const float* __restrict__ z2, + const float* __restrict__ gate_logit, + const float* __restrict__ main_val, + const float* __restrict__ h, + const float* __restrict__ grad_h_carry, + const float* __restrict__ lambda, + int n_batch, + float* __restrict__ grad_w1_scratch, // [B, 5, HEAD_MID, HIDDEN] (+=) + float* __restrict__ grad_b1_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_w2_scratch, // [B, 5, HEAD_MID, HEAD_MID] (+=) + float* __restrict__ grad_b2_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_w_gate_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_b_gate_scratch, // [B, 5] (+=) + float* __restrict__ grad_w_main_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_b_main_scratch, // [B, 5] (+=) + float* __restrict__ grad_w_skip_scratch, // [B, 5, HIDDEN] (+=) + float* __restrict__ grad_b_skip_scratch, // [B, 5] (+=) + float* __restrict__ grad_h // [B, HIDDEN] overwrite (per-batch) +); +``` + +Body: take the existing kernel, replace the `for (int bi ...)` outer loop with `int bi = blockIdx.x; int tid = threadIdx.x; if (bi >= n_batch || tid >= HEAD_MID_H) return;`, and replace every `grad_X[...]` write with `grad_X_scratch[(bi * ...) + ...]`. + +- [ ] **Step 3: Compile to verify the cubin builds** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo build -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build (Rust side will fail on the trainer call until Task 12 lands — that's expected here). + +### Task 12: Wire GRN scratch buffers + reducer launches into the trainer + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/perception.rs` + +- [ ] **Step 1: Add 10 GRN scratch fields to `PerceptionTrainer` struct** + +In the same struct block where you added cfc scratch fields (Task 4), add: + +```rust + // GRN per-batch grad scratch (Phase B). One scratch per param tensor. + grn_grad_w1_scratch_d: CudaSlice, // [B, 5, HEAD_MID, HIDDEN] + grn_grad_b1_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_w2_scratch_d: CudaSlice, // [B, 5, HEAD_MID, HEAD_MID] + grn_grad_b2_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_w_gate_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_b_gate_scratch_d: CudaSlice, // [B, 5] + grn_grad_w_main_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_b_main_scratch_d: CudaSlice, // [B, 5] + grn_grad_w_skip_scratch_d: CudaSlice, // [B, 5, HIDDEN] + grn_grad_b_skip_scratch_d: CudaSlice, // [B, 5] +``` + +- [ ] **Step 2: Allocate the scratch buffers in `PerceptionTrainer::new`** + +Find the existing GRN init block (search for `// ── TFT GRN heads init ──`). After the existing AdamW init for GRN, add: + +```rust + let grn_grad_w1_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HIDDEN_DIM)?; + let grn_grad_b1_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_w2_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?; + let grn_grad_b2_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_w_gate_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_b_gate_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; + let grn_grad_w_main_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_b_main_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; + let grn_grad_w_skip_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HIDDEN_DIM)?; + let grn_grad_b_skip_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; +``` + +- [ ] **Step 3: Populate the new fields in `Ok(Self { ... })`** + +Add them in the struct init block: + +```rust + grn_grad_w1_scratch_d, + grn_grad_b1_scratch_d, + grn_grad_w2_scratch_d, + grn_grad_b2_scratch_d, + grn_grad_w_gate_scratch_d, + grn_grad_b_gate_scratch_d, + grn_grad_w_main_scratch_d, + grn_grad_b_main_scratch_d, + grn_grad_w_skip_scratch_d, + grn_grad_b_skip_scratch_d, +``` + +- [ ] **Step 4: Replace GRN-grad memset_zeros with scratch memset_zeros** + +In `dispatch_train_step`, find the GRN grad zeroing block (search for `// GRN heads: 10 grad accumulators`). Replace the 10 `memset_zeros(&mut self.grad_heads_*_d)` calls with the scratch equivalents: + +```rust + // GRN per-batch grad scratch: zero ONCE per step; K-loop bwd + // accumulates into these, then reduce_axis0 collapses → final + // grad buffers (OVERWRITE) after the K-loop. + self.stream.memset_zeros(&mut self.grn_grad_w1_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w1_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b1_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b1_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w2_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w2_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b2_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b2_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_gate_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_gate_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_gate_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_gate_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_main_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_main_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_main_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_main_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_skip_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_skip_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_skip_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_skip_scratch: {e}"))?; +``` + +DELETE the previous `grad_heads_*_d` zeroings. + +- [ ] **Step 5: Update GRN backward kernel call** + +Find the GRN backward launch inside the K-loop (search for `heads_grn_bwd_fn`). Replace the `.arg(...)` chain with: + +```rust + unsafe { + let mut launch = self.stream.launch_builder(&self.heads_grn_bwd_fn); + launch + .arg(&self.heads_w1_d).arg(&self.heads_w2_d) + .arg(&self.heads_w_gate_d).arg(&self.heads_w_main_d) + .arg(&self.heads_w_skip_d) + .arg(&probs_k_ptr).arg(&gprobs_k_ptr) + .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(&self.lambda_d) + .arg(&n_batch_i) + .arg(&mut self.grn_grad_w1_scratch_d) + .arg(&mut self.grn_grad_b1_scratch_d) + .arg(&mut self.grn_grad_w2_scratch_d) + .arg(&mut self.grn_grad_b2_scratch_d) + .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")?; + } +``` + +Also update the launch config for `cfg_grn_bwd`: it was `grid=(1,1,1)`. Change to: + +```rust + let cfg_grn_bwd = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (HEAD_MID_DIM as u32, 1, 1), + shared_mem_bytes: 0, + }; +``` + +- [ ] **Step 6: Add 10 reducer launches after the K-loop** + +In the same block where you added the cfc reducer launches (Task 5, Step 5), add a GRN reducer block right after (use the same `reduce_at` closure): + +```rust + // ── 8d. Reduce GRN per-batch grad scratch → final grad buffers. + { + let n_batch_i = b_sz as i32; + let reduce_block: u32 = 256; + let reduce_at = |n_tail: usize, + scratch: &CudaSlice, + out: &mut CudaSlice, + label: &'static str| + -> Result<()> { + let cfg = LaunchConfig { + grid_dim: (n_tail as u32, 1, 1), + block_dim: (reduce_block, 1, 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(()) + }; + let h = HIDDEN_DIM; + let m = HEAD_MID_DIM; + let nh = N_HORIZONS; + reduce_at(nh * m * h, &self.grn_grad_w1_scratch_d, + &mut self.grad_heads_w1_d, "reduce grn_grad_w1")?; + reduce_at(nh * m, &self.grn_grad_b1_scratch_d, + &mut self.grad_heads_b1_d, "reduce grn_grad_b1")?; + reduce_at(nh * m * m, &self.grn_grad_w2_scratch_d, + &mut self.grad_heads_w2_d, "reduce grn_grad_w2")?; + reduce_at(nh * m, &self.grn_grad_b2_scratch_d, + &mut self.grad_heads_b2_d, "reduce grn_grad_b2")?; + reduce_at(nh * m, &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, + &mut self.grad_heads_b_gate_d, "reduce grn_grad_b_gate")?; + reduce_at(nh * m, &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, + &mut self.grad_heads_b_main_d, "reduce grn_grad_b_main")?; + reduce_at(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, + &mut self.grad_heads_b_skip_d, "reduce grn_grad_b_skip")?; + } +``` + +- [ ] **Step 7: Build** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo check -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. + +### Task 13: Run smokes for Commit 2 + +- [ ] **Step 1: Full perception_overfit suite** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit -- --nocapture 2>&1 | tail -25 +``` + +Expected: all 10 tests (8 existing + the 2 added in Commit 1) still pass. + +- [ ] **Step 2: B=1 oracle still passes (sanity check)** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test cfc_bwd_b1_oracle -- --nocapture 2>&1 | tail -10 +``` + +Expected: still passes. (Commit 2 didn't touch cfc kernels but cfc oracle is the canary for kernel module loading.) + +### Task 14: Commit Commit 2 + +- [ ] **Step 1: Stage + commit** + +```bash +git add crates/ml-alpha/cuda/multi_horizon_heads.cu \ + crates/ml-alpha/src/trainer/perception.rs +git commit -m "$(cat <<'EOF' +perf(ml-alpha): block-per-batch GRN bwd refactor (K-loop Phase B commit 2) + +multi_horizon_heads_grn_bwd_batched refactored to grid=(B,1,1). Removes +the single-SM bottleneck on the second-most-called K-loop kernel +(64×/step like cfc_bwd). + +Adds 10 per-batch grad scratch buffers (one per GRN param tensor) + +10 reduce_axis0 launches collapsing B → final grad after the K-loop. +~8 MB additional GPU memory at B=32. + +All 10 perception_overfit smokes pass (8 original + the 2 added in +Phase B commit 1). The B=32 convergence smoke validates the new +cross-batch reducer path for the GRN gradients specifically. + +Co-Authored-By: Claude Opus 4.7 +EOF +)" +git push origin ml-alpha-phase-a 2>&1 | tail -3 +``` + +Expected: push succeeds. + +--- + +## Commit 3: VSN bwd refactor + +Same pattern, applied to `variable_selection_bwd`. Two parameter tensors → two scratch buffers → two reducer launches. + +### Task 15: Refactor VSN bwd kernel + +**Files:** +- Modify: `crates/ml-alpha/cuda/variable_selection.cu` + +- [ ] **Step 1: Replace the kernel body** + +Find `extern "C" __global__ void variable_selection_bwd(` and replace its body. Signature change (rename outputs to `_scratch`): + +```cuda +extern "C" __global__ void variable_selection_bwd( + const float* __restrict__ W_vsn, + const float* __restrict__ x, + const float* __restrict__ gates, + const float* __restrict__ grad_y, + int n_rows, + float* __restrict__ grad_W_vsn_scratch, // [n_rows, FEATURE_DIM, FEATURE_DIM] (+=) + float* __restrict__ grad_b_vsn_scratch, // [n_rows, FEATURE_DIM] (+=) + float* __restrict__ grad_x // [n_rows, FEATURE_DIM] overwrite +); +``` + +Body: replace the single-block + internal `for (int row = 0; row < n_rows; ++row)` pattern with `int row = blockIdx.x; int tid = threadIdx.x;`. Each block handles one row. The existing chain rule body stays — only swap the `grad_W_vsn[i * FEATURE_DIM + j] += dl_t * xj` writes for `grad_W_vsn_scratch[((long long)row * FEATURE_DIM + i) * FEATURE_DIM + j] += dl_t * xj`, and similarly for `grad_b_vsn`. + +**Note on n_rows semantics:** VSN's `n_rows = B * K` (one row per (batch, K-position) pair). So scratch is naturally `[B*K, FEATURE_DIM, FEATURE_DIM]` not `[B, ...]`. The reducer treats `n_batch = B*K` and tail dim = `FEATURE_DIM * FEATURE_DIM`. This is the SAME reducer kernel — `reduce_axis0` works on any 2D layout where axis 0 is the to-be-summed dim. + +- [ ] **Step 2: Build to verify the cubin** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo build -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build (Rust side will fail until Task 16). + +### Task 16: Wire VSN scratch + reducer + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/perception.rs` + +- [ ] **Step 1: Add VSN scratch fields** + +```rust + vsn_grad_w_scratch_d: CudaSlice, // [B*K, FEATURE_DIM, FEATURE_DIM] + vsn_grad_b_scratch_d: CudaSlice, // [B*K, FEATURE_DIM] +``` + +- [ ] **Step 2: Allocate in `PerceptionTrainer::new`** + +Find the existing VSN init block. Append: + +```rust + let vsn_grad_w_scratch_d = stream.alloc_zeros::( + cfg.n_batch * cfg.seq_len * FEATURE_DIM * FEATURE_DIM)?; + let vsn_grad_b_scratch_d = stream.alloc_zeros::( + cfg.n_batch * cfg.seq_len * FEATURE_DIM)?; +``` + +- [ ] **Step 3: Populate fields in `Ok(Self { ... })`** + +```rust + vsn_grad_w_scratch_d, + vsn_grad_b_scratch_d, +``` + +- [ ] **Step 4: Replace VSN grad memset_zeros with scratch memset_zeros** + +In `dispatch_train_step`, find the VSN grad zero block. Replace: + +```rust + self.stream.memset_zeros(&mut self.grad_vsn_w_d) + .map_err(|e| anyhow::anyhow!("zero grad_vsn_w: {e}"))?; + self.stream.memset_zeros(&mut self.grad_vsn_b_d) + .map_err(|e| anyhow::anyhow!("zero grad_vsn_b: {e}"))?; +``` + +with: + +```rust + 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}"))?; +``` + +- [ ] **Step 5: Update VSN backward launch + add reducer launches** + +Find the VSN bwd launch (search for `variable_selection_bwd`). The existing call uses `grid=(1,1,1)`. Change to: + +```rust + 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), // VSN_BLOCK + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.vsn_bwd_fn); + launch + .arg(&self.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")?; } +``` + +Then add reducer launches RIGHT AFTER (before `opt_vsn_w.step(...)`): + +```rust + // Reduce VSN per-row grad scratch → final grad buffers. + { + let n_rows_i = (b_sz * k_seq) as i32; + let cfg_red_vsn_w = LaunchConfig { + grid_dim: (FEATURE_DIM as u32 * FEATURE_DIM as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }; + let n_tail_w = (FEATURE_DIM * FEATURE_DIM) as i32; + 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); + unsafe { launch.launch(cfg_red_vsn_w).context("reduce vsn_grad_w")?; } + + let cfg_red_vsn_b = LaunchConfig { + grid_dim: (FEATURE_DIM as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }; + let n_tail_b = FEATURE_DIM as i32; + 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); + unsafe { launch.launch(cfg_red_vsn_b).context("reduce vsn_grad_b")?; } + } +``` + +- [ ] **Step 6: Build** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo check -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. + +### Task 17: Smokes + commit for Commit 3 + +- [ ] **Step 1: Run perception_overfit** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit -- --nocapture 2>&1 | tail -25 +``` + +Expected: all 10 tests pass. + +- [ ] **Step 2: Commit and push** + +```bash +git add crates/ml-alpha/cuda/variable_selection.cu \ + crates/ml-alpha/src/trainer/perception.rs +git commit -m "$(cat <<'EOF' +perf(ml-alpha): block-per-row VSN bwd refactor (K-loop Phase B commit 3) + +variable_selection_bwd refactored from grid=(1,1,1) to grid=(B*K,1,1). +VSN's n_rows = B*K positions; block-per-row matches the existing fwd +layout (which is already grid=(B*K,1,1)). + +Adds 2 per-row grad scratch buffers + 2 reduce_axis0 launches. ~210 KB +scratch at B=32, K=64. + +VSN bwd is 1×/step (not in K-loop) so the wall-time win here is small +compared to commits 1+2. Done for pattern uniformity — every per-batch +or per-row bwd in the trainer now uses scratch+reducer. + +All 10 perception_overfit smokes pass. + +Co-Authored-By: Claude Opus 4.7 +EOF +)" +git push origin ml-alpha-phase-a 2>&1 | tail -3 +``` + +Expected: push succeeds. + +--- + +## Commit 4: attention_pool bwd refactor + +Same pattern, applied to `attention_pool_bwd`. One parameter tensor (Q) → one scratch buffer → one reducer launch. + +### Task 18: Refactor attention_pool bwd kernel + +**Files:** +- Modify: `crates/ml-alpha/cuda/attention_pool.cu` + +- [ ] **Step 1: Replace the bwd kernel signature + body** + +Find `extern "C" __global__ void attention_pool_bwd(`. Signature change: + +```cuda +extern "C" __global__ void attention_pool_bwd( + const float* __restrict__ Q, + const float* __restrict__ ln_out, + const float* __restrict__ attn_weights, + const float* __restrict__ grad_context, + int n_batch, + int k_seq, + float* __restrict__ grad_Q_scratch, // [B, HIDDEN] (+=) + float* __restrict__ grad_ln_out // [B, K, HIDDEN] += (per-batch indexed) +); +``` + +Body: replace `for (int bi = 0; bi < n_batch; ++bi)` outer loop with `int bi = blockIdx.x;`. Each block handles one batch. The `grad_Q[tid] += dq_local` write becomes `grad_Q_scratch[(long long)bi * ATTN_HIDDEN_DIM + tid] += dq_local`. The `grad_ln_b[k * ATTN_HIDDEN_DIM + tid] += ...` write is already per-batch indexed (`ln_b` is `grad_ln_out + bi * k_seq * HIDDEN`) so no change needed — single writer per (bi, k, tid). + +- [ ] **Step 2: Build to verify** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo build -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. + +### Task 19: Wire attention pool scratch + reducer + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/perception.rs` + +- [ ] **Step 1: Add scratch field** + +```rust + attn_grad_q_scratch_d: CudaSlice, // [B, HIDDEN] +``` + +- [ ] **Step 2: Allocate** + +In the existing attention pool init block in `PerceptionTrainer::new`: + +```rust + let attn_grad_q_scratch_d = stream.alloc_zeros::( + cfg.n_batch * HIDDEN_DIM)?; +``` + +- [ ] **Step 3: Populate field** + +```rust + attn_grad_q_scratch_d, +``` + +- [ ] **Step 4: Replace attn grad zero** + +Find `self.stream.memset_zeros(&mut self.grad_attn_q_d)`. Replace with: + +```rust + self.stream.memset_zeros(&mut self.attn_grad_q_scratch_d) + .map_err(|e| anyhow::anyhow!("zero attn_grad_q_scratch: {e}"))?; +``` + +- [ ] **Step 5: Update attention_pool_bwd launch + add reducer** + +Find the attention_pool_bwd launch (search for `attn_bwd_fn`). Change launch config to `grid=(b_sz,1,1)`: + +```rust + 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), + shared_mem_bytes: shared as u32, + }; + let mut launch = self.stream.launch_builder(&self.attn_bwd_fn); + launch + .arg(&self.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")?; } +``` + +Then add reducer right after (before `opt_attn_q.step(...)`): + +```rust + // Reduce attn per-batch grad_Q scratch → final grad_attn_q_d. + { + let n_batch_i = b_sz as i32; + let n_tail_i = HIDDEN_DIM as i32; + let cfg_red = LaunchConfig { + grid_dim: (HIDDEN_DIM as u32, 1, 1), + block_dim: (256, 1, 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")?; } + } +``` + +- [ ] **Step 6: Build** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo check -p ml-alpha --features cuda 2>&1 | tail -10 +``` + +Expected: clean build. + +### Task 20: Smokes + commit for Commit 4 + +- [ ] **Step 1: Run perception_overfit** + +```bash +CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true cargo test -p ml-alpha --features cuda --test perception_overfit -- --nocapture 2>&1 | tail -25 +``` + +Expected: all 10 tests pass. + +- [ ] **Step 2: Commit and push** + +```bash +git add crates/ml-alpha/cuda/attention_pool.cu \ + crates/ml-alpha/src/trainer/perception.rs +git commit -m "$(cat <<'EOF' +perf(ml-alpha): block-per-batch attention pool bwd refactor (K-loop Phase B commit 4) + +attention_pool_bwd refactored from grid=(1,1,1) to grid=(B,1,1). The +existing per-batch grad_ln_out writes were already uniquely indexed; +only grad_Q needed scratch+reducer (1 scratch, 1 reducer launch). + +~16 KB scratch at B=32 — trivially small. + +attn_pool bwd is 1×/step (not in K-loop) so the wall-time win is tiny — +this commit is for pattern uniformity. With this, every single-SM bwd +kernel in the trainer has been refactored. + +All 10 perception_overfit smokes pass. + +Phase B kernel work complete. Next: local + cluster A/B perf +benchmark to verify gates #6, #7, #8 from the spec. + +Co-Authored-By: Claude Opus 4.7 +EOF +)" +git push origin ml-alpha-phase-a 2>&1 | tail -3 +``` + +Expected: push succeeds. + +--- + +## Task 21: Local A/B performance benchmark + +**Files:** none modified. Diagnostic-only. + +- [ ] **Step 1: Capture pre-fix baseline (one-time)** + +If you don't already have a release-build wall time for the B=32 synthetic smoke at the pre-Phase-B commit (`7a558b88b`), check out that commit briefly, build release, time the smoke, then check back out: + +```bash +git stash +git checkout 7a558b88b +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 \ + cargo build --release -p ml-alpha --features cuda 2>&1 | tail -3 +time (CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true \ + cargo test --release -p ml-alpha --features cuda --test perception_overfit \ + stacked_trainer_loss_shrinks_at_batch_32 2>&1 | tail -5) +# Note: this test doesn't exist on 7a558b88b — substitute +# stacked_trainer_loss_shrinks_on_constant_signal (B=1) and DOUBLE +# the wall time as a B=32 estimate. Or skip this step and rely on +# cluster A/B (Task 22) for the real measurement. +git checkout ml-alpha-phase-a +git stash pop +``` + +If the pre-fix B=32 test doesn't exist (it doesn't — we added it in Phase B commit 1), skip directly to cluster A/B for the real performance measurement. + +- [ ] **Step 2: Run post-fix wall time** + +```bash +SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 \ + cargo build --release -p ml-alpha --features cuda 2>&1 | tail -3 +time (CUDA_COMPUTE_CAP=86 SQLX_OFFLINE=true \ + cargo test --release -p ml-alpha --features cuda --test perception_overfit \ + stacked_trainer_loss_shrinks_at_batch_32 2>&1 | tail -5) +``` + +Record wall time. Pre-fix expectation per spec: ~250 ms/step at B=32 (interpolated from t6z89's 640 ms/step at the production config). Post-fix target: ≤ 80 ms/step (~3× speedup). + +This local measurement is informational; the authoritative perf gate is the cluster A/B (Task 22). + +## Task 22: Cluster A/B performance benchmark + +**Files:** none modified. Diagnostic-only. + +- [ ] **Step 1: Submit the post-fix run** + +```bash +./scripts/argo-alpha-perception.sh \ + --branch ml-alpha-phase-a --sha "$(git rev-parse HEAD)" \ + --batch-size 32 --auto-horizon-weights \ + --mamba2-state-dim 32 --seq-len 64 \ + --decision-stride 4 \ + --gpu-pool ci-training-l40s \ + --early-stop-metric auc_h6000 --early-stop-patience 5 \ + --epochs 5 --n-train-seqs 24000 \ + --cv-fold 1 --cv-n-folds 3 --cv-train-window 4 2>&1 | tail -10 +``` + +(5 epochs is enough to measure per-epoch wall time. We don't need the full 30 — the comparison is t6z89's per-epoch wall.) + +Note the new workflow ID from the output (e.g., `alpha-perception-XXXXX`). + +- [ ] **Step 2: Monitor + record per-epoch wall time** + +Wait for the workflow to complete the first 3-5 epochs. Then collect timestamps: + +```bash +kubectl logs -n foxhunt alpha-perception-XXXXX-train- -c main 2>&1 \ + | grep -E "epoch complete" \ + | awk -F'epoch complete epoch=' '{print $1, $2}' \ + | head -10 +``` + +Compute per-epoch wall time (delta between consecutive `epoch complete` timestamps). + +**Pre-fix baseline (t6z89):** 8 min / epoch. +**Acceptance gate #7:** ≤ 2.5 min / epoch (≥ 3× speedup). Stretch ≤ 1.5 min (≥ 5×). + +- [ ] **Step 3: Compare AUC trajectory to t6z89** + +Same log scrape: + +```bash +kubectl logs -n foxhunt alpha-perception-XXXXX-train- -c main 2>&1 \ + | grep -E "validation epoch=" \ + | awk -F'auc_' '{print $1, $2}' \ + | head -10 +``` + +**Acceptance gate #8:** at each epoch N, the new run's `mean_auc` and `auc_h6000` are within ±0.005 of t6z89's values at the same epoch N. If deviation exceeds ±0.005, the refactor changed semantics — investigate (likely the scratch accumulation order or the reducer's sum order is producing different gradients than the old sequential path expected, beyond what FP non-associativity alone explains). + +t6z89 per-epoch reference: +| Epoch | mean_auc | auc_h6000 | +|-------|----------|-----------| +| 0 | 0.7428 | 0.7211 | +| 1 | 0.7465 | 0.7359 | +| 2 | 0.7472 | n/a (val_loss=0.5728) | +| 3 | 0.7470 | 0.7434 | + +(Later epochs as t6z89 progresses — pull from the bsml6 monitor logs.) + +- [ ] **Step 4: Final commit (only on gate pass)** + +If gates #7 and #8 both pass, no further commits — Phase B is complete and the branch is ready for downstream work (3-fold CV at the new speed). + +If gate #7 fails (perf shortfall): the work is salvageable; profile dispatch_train_step to find the next bottleneck (Mamba2 scan or cuBLAS). Don't revert. Document the achieved speedup and where the floor is. + +If gate #8 fails (AUC deviation): execute the rollback plan from the spec — `git revert` the commit(s) that introduced the deviation. Likely candidates in order: Commit 1 (cfc bwd is the most numerically sensitive). Re-run cluster A/B after each revert to localize. + +--- + +## Self-Review + +**Spec coverage check:** + +- Spec § "Architecture Choice" (block-per-batch + scratch + reducer) → Tasks 1, 2, 3, 11, 15, 18 (kernel refactors). ✓ +- Spec § "Per-K-iteration accumulation invariant" → Steps in Tasks 5, 12, 16, 19 (memset scratch once at step start, += in K-loop, reducer once after K-loop). ✓ +- Spec § "AdamW-after-reducer invariant" → Documented in Commit 1's commit message + the trainer scratch comments in Task 4 Step 2. ✓ +- Spec § "Per-Kernel Breakdown" table (5 kernels, 17 reducer launches) → Tasks cover all 5 kernels; reducer launch counts: cfc=4 (Task 5), grn=10 (Task 12), vsn=2 (Task 16), attn=1 (Task 19) = 17. ✓ +- Spec § "Trainer Wiring" → Tasks 4, 5, 12, 16, 19 add 17 new fields + memsets + reducer launches as specified. ✓ +- Spec § "Build System" → Task 1 Step 2 (KERNELS + cache-bust v11). ✓ +- Spec § "Tests + acceptance" gate #1 (perception_overfit) → Tasks 9, 13, 17, 20. ✓ +- Spec § gate #2 (B=32 smoke) → Task 7. ✓ +- Spec § gate #3 (backward_finite_diff) → Task 9 Step 2. ✓ +- Spec § gate #4 (cfc B=1 oracle) → Task 6. ✓ +- Spec § gate #5 (scratch-clears) → Task 8. ✓ +- Spec § gate #6 (local perf) → Task 21. ✓ +- Spec § gate #7 (cluster A/B wall) → Task 22 Step 2. ✓ +- Spec § gate #8 (AUC trajectory) → Task 22 Step 3. ✓ +- Spec § "Rollout" 4-commit structure → Commit 1 (Tasks 1-10), Commit 2 (Tasks 11-14), Commit 3 (Tasks 15-17), Commit 4 (Tasks 18-20). ✓ +- Spec § "Rollback" → Task 22 Step 4 references the spec rollback. ✓ + +**Placeholder scan:** + +- No "TBD", "TODO", "fill in details" tokens. +- One callout in Task 11 Step 2: "Mechanical translation rules above. Reference the existing kernel for the chain-rule math; only the indexing and launch contract change." — this is intentional. The chain rule body is ~150 lines and identical to the existing kernel; duplicating it verbatim in the plan would be a code-listing not a plan. The mechanical translation rules ARE specified in detail; the kernel just inherits the existing math. +- One callout in Task 6 Step 1: "If that exact API name doesn't exist, fall back to..." — the cudarc 0.19 upload API is `stream.memcpy_stod` per recent usage; the fallback is defensive in case the implementer hits a version mismatch. + +**Type consistency:** + +- Field names: `cfc_grad_w_in_scratch_d` (Task 4) → used in Tasks 5, 9, 10. ✓ +- `reduce_axis0_fn` (Task 4) → used in Tasks 5, 12, 16, 19. ✓ +- `grn_grad_w1_scratch_d` etc. (Task 12) consistent across the GRN block. ✓ +- VSN scratch fields (`vsn_grad_w_scratch_d`, `vsn_grad_b_scratch_d`) consistent across Task 16. ✓ +- Attn scratch (`attn_grad_q_scratch_d`) consistent in Task 19. ✓ +- Kernel signature: `reduce_axis0(per_batch, n_batch, n_tail, out)` (Task 1) consistent at every call site (Tasks 5, 12, 16, 19). ✓ + +No issues found in self-review. + +--- + +## Execution Handoff + +Plan complete and saved to `docs/superpowers/plans/2026-05-17-kloop-parallelization.md`. Two execution options: + +**1. Subagent-Driven (recommended)** — I dispatch a fresh subagent per task, review between tasks, fast iteration + +**2. Inline Execution** — Execute tasks in this session using executing-plans, batch execution with checkpoints + +**Which approach?**