Files
foxhunt/docs/superpowers/plans/2026-05-17-kloop-parallelization.md
jgrusewski c508d101c1 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 <noreply@anthropic.com>
2026-05-17 23:27:20 +02:00

1862 lines
78 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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<f32>, // [B, n_hid, n_in]
cfc_grad_w_rec_scratch_d: CudaSlice<f32>, // [B, n_hid, n_hid]
cfc_grad_b_scratch_d: CudaSlice<f32>, // [B, n_hid]
cfc_grad_tau_scratch_d: CudaSlice<f32>, // [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<CudaModule>,
```
- [ ] **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::<f32>(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::<f32>(cfg.n_batch * n_hid * n_hid)
.context("cfc_grad_w_rec_scratch_d alloc")?;
let cfc_grad_b_scratch_d = stream
.alloc_zeros::<f32>(cfg.n_batch * n_hid)
.context("cfc_grad_b_scratch_d alloc")?;
let cfc_grad_tau_scratch_d = stream
.alloc_zeros::<f32>(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::<f32>()) 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<f32>,
out: &mut CudaSlice<f32>,
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<f32> = (0..n_in).map(|_| r.gen_range(-1.0..1.0)).collect();
let h_old: Vec<f32> = (0..n_hid).map(|_| r.gen_range(-1.0..1.0)).collect();
let grad_h_new: Vec<f32> = (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<f32> {
let h = cudarc::driver::result::malloc_host::<f32>(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::<f32>(n_batch * n_hid * n_in).expect("scratch");
let mut grad_w_rec_scratch_d = stream.alloc_zeros::<f32>(n_batch * n_hid * n_hid).expect("scratch");
let mut grad_b_scratch_d = stream.alloc_zeros::<f32>(n_batch * n_hid).expect("scratch");
let mut grad_tau_scratch_d = stream.alloc_zeros::<f32>(n_batch * n_hid).expect("scratch");
let mut grad_h_old_d = stream.alloc_zeros::<f32>(n_batch * n_hid).expect("grad_h_old");
let mut grad_x_d = stream.alloc_zeros::<f32>(n_batch * n_in).expect("grad_x");
let smem = (2 * n_hid * std::mem::size_of::<f32>()) 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::<f32>(n_hid * n_in).expect("final");
let mut grad_w_rec_d = stream.alloc_zeros::<f32>(n_hid * n_hid).expect("final");
let mut grad_b_d = stream.alloc_zeros::<f32>(n_hid).expect("final");
let mut grad_tau_d = stream.alloc_zeros::<f32>(n_hid).expect("final");
let reduce = |scratch: &cudarc::driver::CudaSlice<f32>,
out: &mut cudarc::driver::CudaSlice<f32>,
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<f32>` 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<Mbp10RawInput>> = Vec::with_capacity(cfg.n_batch);
let mut warm_labels: Vec<Vec<[f32; 5]>> = 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<Mbp10RawInput>> = Vec::with_capacity(cfg.n_batch);
let mut labels: Vec<Vec<[f32; 5]>> = 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<Mbp10RawInput>> = Vec::with_capacity(cfg.n_batch);
let mut labels: Vec<Vec<[f32; 5]>> = 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<Vec<f32>> {
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::<f32>(),
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 <noreply@anthropic.com>
EOF
)"
```
Expected output: `[ml-alpha-phase-a <sha>] 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<f32>, // [B, 5, HEAD_MID, HIDDEN]
grn_grad_b1_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_w2_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID, HEAD_MID]
grn_grad_b2_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_w_gate_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_b_gate_scratch_d: CudaSlice<f32>, // [B, 5]
grn_grad_w_main_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_b_main_scratch_d: CudaSlice<f32>, // [B, 5]
grn_grad_w_skip_scratch_d: CudaSlice<f32>, // [B, 5, HIDDEN]
grn_grad_b_skip_scratch_d: CudaSlice<f32>, // [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::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HIDDEN_DIM)?;
let grn_grad_b1_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_w2_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?;
let grn_grad_b2_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_w_gate_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_b_gate_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS)?;
let grn_grad_w_main_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_b_main_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS)?;
let grn_grad_w_skip_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HIDDEN_DIM)?;
let grn_grad_b_skip_scratch_d = stream.alloc_zeros::<f32>(
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<f32>,
out: &mut CudaSlice<f32>,
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 <noreply@anthropic.com>
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<f32>, // [B*K, FEATURE_DIM, FEATURE_DIM]
vsn_grad_b_scratch_d: CudaSlice<f32>, // [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::<f32>(
cfg.n_batch * cfg.seq_len * FEATURE_DIM * FEATURE_DIM)?;
let vsn_grad_b_scratch_d = stream.alloc_zeros::<f32>(
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 <noreply@anthropic.com>
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<f32>, // [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::<f32>(
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::<f32>();
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 <noreply@anthropic.com>
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-<hash> -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-<hash> -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?**