# Custom Kernel Optimization 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:** Rewrite 7 custom CUDA kernels that consume 95% of GPU time — replace inner matmul loops with cuBLAS GEMMs, replace batch-serial patterns with batch-parallel grids. Target: 1950ms/step → <50ms/step. **Architecture:** Split each "one-thread-per-weight-loops-over-batch" kernel into: (1) cuBLAS GEMMs for matrix projections (tensor cores), (2) lightweight element-wise/scan kernels for non-GEMM ops. The key insight from NVIDIA's official Mamba implementation: `grid=(batch, dim)` not `grid=(params)`. **Tech Stack:** CUDA 13, cublasLtMatmul (TF32), cudarc (vendored), existing IqnGemmDesc pattern for cached GEMM descriptors. **Dimensions:** SH2=256, STATE_D=16, K=8, batch=8192, IQN_Q=64, adv_h=128. --- ### Task 1: Mamba2 Forward — cuBLAS Projections + Lightweight Scan (1.1%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu` - Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` The forward scan currently does per-sample inner matmuls: `for j in 0..SH2: a_val += W_A[j,s] * h_t[j]`. This is a GEMM `[STATE_D, SH2] × [SH2, B]` that should use cuBLAS. - [ ] **Step 1: Allocate scratch buffers for pre-projected values** In `gpu_dqn_trainer.rs` constructor, allocate: ```rust // Mamba2 pre-projection scratch: A_vals[B*K, STATE_D], B_vals[B*K, STATE_D] let mamba2_a_proj = alloc_f32(&stream, batch_size * MAMBA2_HISTORY_K * MAMBA2_STATE_DIM, "mamba2_a_proj")?; let mamba2_b_proj = alloc_f32(&stream, batch_size * MAMBA2_HISTORY_K * MAMBA2_STATE_DIM, "mamba2_b_proj")?; let mamba2_c_proj = alloc_f32(&stream, batch_size * MAMBA2_STATE_DIM, "mamba2_c_proj")?; ``` - [ ] **Step 2: Create cuBLAS GEMM descriptors for Mamba2 projections** Using the existing `create_cached_fwd_gemm_desc` pattern, create 3 GEMM descriptors: ``` // A_proj = W_A^T @ h_history_flat: m=STATE_D(16), n=B*K(65536), k=SH2(256) // B_proj = W_B^T @ h_history_flat: same dims // C_proj = W_C^T @ h_s2: m=STATE_D(16), n=B(8192), k=SH2(256) ``` Load these from a new CUmodule (graph-captured, must be isolated). - [ ] **Step 3: Write lightweight scan kernel `mamba2_scan_projected`** Replace `mamba2_temporal_scan` with a kernel that reads pre-projected values: ```c extern "C" __global__ void mamba2_scan_projected( const float* __restrict__ a_proj, // [B*K, STATE_D] pre-projected const float* __restrict__ b_proj, // [B*K, STATE_D] pre-projected const float* __restrict__ c_proj, // [B, STATE_D] pre-projected const float* __restrict__ h_s2, // [B, SH2] current trunk float* __restrict__ h_enriched, // [B, SH2] output int B, int K, int sh2, int state_d, const float* __restrict__ isv_signals, const float* __restrict__ temporal_weight) { int b = blockIdx.x; int j = blockIdx.y * blockDim.x + threadIdx.x; if (b >= B || j >= sh2) return; // Scan over K=8 timesteps — reads from pre-projected buffers float context = 0.0f; for (int s = 0; s < state_d; s++) { float x = 0.0f; for (int t = 0; t < K; t++) { int idx = (b * K + t) * state_d + s; float gate = 1.0f / (1.0f + expf(-a_proj[idx])); if (isv_signals) gate *= isv_signals[11]; x = gate * x + b_proj[idx]; } context += c_proj[b * state_d + s] * x; } // Apply temporal routing and residual float h_val = h_s2[b * sh2 + j]; float weight = temporal_weight ? temporal_weight[j] : 1.0f; float blended = h_val + weight * context * (1.0f / (float)state_d); h_enriched[b * sh2 + j] = blended; } ``` Grid: `(B, ceil(SH2/256))`, Block: `(256)`. Each thread handles one (sample, feature). K=8 inner loop — trivial. STATE_D=16 inner loop — 16 iterations, no memory bottleneck. - [ ] **Step 4: Update `mamba2_forward` in gpu_dqn_trainer.rs** Replace the old kernel launch with: 1. cuBLAS GEMM: `W_A^T @ h_history_flat → mamba2_a_proj` 2. cuBLAS GEMM: `W_B^T @ h_history_flat → mamba2_b_proj` 3. cuBLAS GEMM: `W_C^T @ h_s2 → mamba2_c_proj` 4. Launch `mamba2_scan_projected` kernel Note: `h_history` is `[B, K, SH2]` — reshape to `[B*K, SH2]` for the GEMM by treating it as a flat contiguous buffer (it already is in memory). - [ ] **Step 5: Build and test** ```bash SQLX_OFFLINE=true cargo check -p ml SQLX_OFFLINE=true cargo test -p ml --lib ``` - [ ] **Step 6: Commit** ```bash git commit -m "perf: mamba2 forward — cuBLAS projections + lightweight scan (22ms → <1ms)" ``` --- ### Task 2: Mamba2 Backward — cuBLAS Projections + Lightweight Reverse Scan (75%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/experience_kernels.cu` - Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` This is THE bottleneck — 1567ms/call, 75% of total GPU time. Same approach as forward but with reverse scan + weight gradient accumulation via cuBLAS. - [ ] **Step 1: Allocate backward scratch buffers** ```rust // Backward needs: d_gate[B*K, STATE_D], d_x[B*K, STATE_D], d_context[B, STATE_D] let mamba2_d_gate = alloc_f32(&stream, batch_size * MAMBA2_HISTORY_K * MAMBA2_STATE_DIM, "mamba2_d_gate")?; let mamba2_d_x = alloc_f32(&stream, batch_size * MAMBA2_HISTORY_K * MAMBA2_STATE_DIM, "mamba2_d_x")?; let mamba2_d_context = alloc_f32(&stream, batch_size * MAMBA2_STATE_DIM, "mamba2_d_context")?; ``` - [ ] **Step 2: Write `mamba2_scan_backward_projected` kernel** This kernel does the reverse scan using pre-projected values. No inner matmuls. Grid: `(B, STATE_D)`, each thread handles one (sample, state_dim) pair. ```c extern "C" __global__ void mamba2_scan_backward_projected( const float* __restrict__ a_proj, // [B*K, STATE_D] forward A projections const float* __restrict__ b_proj, // [B*K, STATE_D] forward B projections const float* __restrict__ c_proj, // [B, STATE_D] forward C projections const float* __restrict__ d_h_enriched,// [B, SH2] upstream gradient const float* __restrict__ w_c, // [SH2, STATE_D] float* __restrict__ d_gate, // [B*K, STATE_D] output: sigmoid_deriv * d_x * x_prev float* __restrict__ d_x_out, // [B*K, STATE_D] output: d_x per timestep float* __restrict__ d_context, // [B, STATE_D] output: d_context = d_out @ W_C int B, int K, int sh2, int state_d, const float* __restrict__ isv_signals) { int b = blockIdx.x; int s = blockIdx.y * blockDim.x + threadIdx.x; if (b >= B || s >= state_d) return; // Forward replay: recompute x_states[0..K] for this (b, s) float x_states[9]; // K+1 max x_states[0] = 0.0f; float a_raw[8]; for (int t = 0; t < K; t++) { int idx = (b * K + t) * state_d + s; a_raw[t] = a_proj[idx]; float gate = 1.0f / (1.0f + expf(-a_raw[t])); if (isv_signals) gate *= isv_signals[11]; x_states[t + 1] = gate * x_states[t] + b_proj[idx]; } // d_context[b, s] = sum_j(d_h_enriched[b, j] * w_c[j, s]) float d_ctx = 0.0f; for (int j = 0; j < sh2; j++) { d_ctx += d_h_enriched[b * sh2 + j] * w_c[j * state_d + s]; } d_context[b * state_d + s] = d_ctx; // d_x[K] = d_context * c_proj was the last step contribution float d_x_s = d_ctx; // Reverse scan for (int t = K - 1; t >= 0; t--) { int idx = (b * K + t) * state_d + s; float a_val = a_raw[t]; float gate = 1.0f / (1.0f + expf(-a_val)); float sigmoid_deriv = gate * (1.0f - gate); float stability = isv_signals ? isv_signals[11] : 1.0f; // d_gate for W_A gradient: d_x * x_prev * sigmoid' d_gate[idx] = d_x_s * x_states[t] * sigmoid_deriv * stability; // d_x for W_B gradient d_x_out[idx] = d_x_s; // Propagate d_x backward d_x_s = d_x_s * gate * stability; } } ``` - [ ] **Step 3: Create cuBLAS GEMM descriptors for weight gradient accumulation** After the scan kernel produces `d_gate[B*K, STATE_D]` and `d_x[B*K, STATE_D]`: ``` d_W_A[SH2, STATE_D] = h_history^T[SH2, B*K] @ d_gate[B*K, STATE_D] GEMM: transa=N, transb=N, m=STATE_D(16), n=SH2(256), k=B*K(65536) d_W_B[SH2, STATE_D] = h_history^T[SH2, B*K] @ d_x[B*K, STATE_D] same dims d_W_C[SH2, STATE_D] = h_s2^T[SH2, B] @ d_context[B, STATE_D] GEMM: m=STATE_D(16), n=SH2(256), k=B(8192) ``` - [ ] **Step 4: Update `mamba2_backward` in gpu_dqn_trainer.rs** Replace old kernel launch with: 1. Reuse forward A/B projections (already computed in forward, save them) 2. Launch `mamba2_scan_backward_projected` kernel 3. cuBLAS GEMM: `h_history^T @ d_gate → d_W_A` 4. cuBLAS GEMM: `h_history^T @ d_x → d_W_B` 5. cuBLAS GEMM: `h_s2^T @ d_context → d_W_C` - [ ] **Step 5: Build and test** ```bash SQLX_OFFLINE=true cargo check -p ml SQLX_OFFLINE=true cargo test -p ml --lib ``` - [ ] **Step 6: Run local smoke test** ```bash SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib --release --features ml/cuda -- smoke_tests --ignored --nocapture ``` - [ ] **Step 7: Commit** ```bash git commit -m "perf: mamba2 backward — cuBLAS projections + reverse scan (1567ms → <5ms)" ``` --- ### Task 3: Q-Denoise Backward — cuBLAS GEMMs (7.3%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` Replace the custom `q_denoise_backward` kernel (2-layer MLP backward, one thread per weight, loops over batch) with 2× cuBLAS backward GEMMs + bias_grad_reduce. The MLP is 12→24→12 with batch=8192. - [ ] **Step 1: Create cuBLAS GEMM descriptors for denoise backward** ```rust // Layer 2 backward dX: m=24, n=8192, k=12 (dX = W2^T @ dY2) // Layer 2 backward dW: m=12, n=24, k=8192 (dW = dY2^T @ h1) // Layer 1 backward dX: m=24, n=8192, k=24 (dX = W1^T @ dY1) -- not needed (no upstream) // Layer 1 backward dW: m=24, n=24, k=8192 (dW = dY1^T @ Q_input) ``` - [ ] **Step 2: Rewrite `launch_q_denoise_backward` to use cuBLAS** Replace single kernel with: loss gradient → cuBLAS dW2 → ReLU mask → cuBLAS dW1 → bias_grad_reduce × 2. - [ ] **Step 3: Build, test, commit** ```bash git commit -m "perf: q_denoise_backward — cuBLAS GEMMs (152ms → <1ms)" ``` --- ### Task 4: IQN Bias Grad Reduce — Warp-Parallel Reduction (3.9%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu` Current: one thread per output neuron (out_dim=3), each loops over B×Q=524,288. 3 threads doing 524K iterations. - [ ] **Step 1: Rewrite `iqn_bias_grad_reduce` with batch-parallel grid** ```c // Phase 1: each block reduces 256 elements // Grid: (ceil(B*Q/256), out_dim), Block: 256 // Each thread reads one element, warp shuffle reduces to partial sum // Block-level reduction via shared memory → one partial per block // Phase 2: reduce block partials → final db[out_dim] // Grid: (out_dim), Block: 256 // Each thread reads one block partial, warp reduce to final sum ``` - [ ] **Step 2: Build, test, commit** ```bash git commit -m "perf: iqn_bias_grad_reduce — warp-parallel (16ms×130 → 0.1ms×130)" ``` --- ### Task 5: KAN Grad Reduce — Batch-Parallel Reduction (3.2%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/experience_kernels.cu` Same anti-pattern as IQN — one thread per param, loops over batch. Fix: batch-parallel 2-phase reduction. - [ ] **Step 1: Rewrite `kan_grad_reduce` with batch-parallel grid** Same pattern as Task 4 — phase 1 block reduce + phase 2 final reduce. - [ ] **Step 2: Build, test, commit** ```bash git commit -m "perf: kan_grad_reduce — batch-parallel (8ms×208 → 0.1ms×208)" ``` --- ### Task 6: Curiosity Fwd+Bwd — cuBLAS GEMMs (2.1%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/curiosity_training_kernel.cu` - Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` The training kernel is a single launch doing forward + backward for the entire batch. Replace with cuBLAS GEMMs for the curiosity MLP. - [ ] **Step 1: Replace `curiosity_fwd_bwd_per_block` with cuBLAS pipeline** Forward: 2 cuBLAS GEMMs (state→hidden→pred). Backward: 2 cuBLAS GEMMs (dW). Loss: element-wise kernel. - [ ] **Step 2: Build, test, commit** ```bash git commit -m "perf: curiosity training — cuBLAS GEMMs (1085ms → <5ms)" ``` --- ### Task 7: Risk Budget Forward — cuBLAS GEMMs (1.0%) **Files:** - Modify: `crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs` - Modify: `crates/ml/src/cuda_pipeline/experience_kernels.cu` Per-sample inner matmul over SH2=256. Replace with cuBLAS. - [ ] **Step 1: Replace `risk_budget_forward` with cuBLAS pipeline** Concat [h_s2; ISV; pred_error] → cuBLAS GEMM1 → ReLU → cuBLAS GEMM2 → sigmoid. Same pattern as main forward. - [ ] **Step 2: Build, test, commit** ```bash git commit -m "perf: risk_budget_forward — cuBLAS GEMMs (20ms → <0.1ms)" ``` --- ### Task 8: Deploy + nsys Validation - [ ] **Step 1: Push and deploy with nsys** ```bash git push origin main ./scripts/argo-train.sh dqn --baseline --epochs 1 --sanitizer nsys ``` - [ ] **Step 2: Verify step timing** Expected: `STEP_TIMING step=50 step_ms=<50` (was 1950ms) - [ ] **Step 3: Compare nsys kernel breakdown** Verify all 7 kernels are now <5ms each. cuBLAS GEMMs should dominate. - [ ] **Step 4: Run smoke tests on H100** All 20 smoke tests pass. Training metrics within 5% of baseline. - [ ] **Step 5: Commit validation results** ```bash git commit -m "docs: kernel optimization validation — step time Xms (was 1950ms)" ```