nsys data: custom kernels = 95% of GPU time. mamba2_scan_backward = 75%. Root cause: one-thread-per-weight anti-pattern serializes B=8192 samples. Fix: cuBLAS for projections, lightweight scans for sequential state. Target: 1950ms/step → <50ms/step. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
6.0 KiB
Custom Kernel Optimization — cuBLAS Projections + Batch-Parallel Scans
Problem
nsys profiling reveals custom CUDA kernels consume 95% of GPU time. cuBLAS GEMMs take <1%. The root cause: all custom kernels use a "one thread per weight, loop over batch" anti-pattern that serializes B=8192 samples into each thread.
nsys Breakdown (batch=8192, H100 PCIe)
| Kernel | Time% | ms/call | Calls | Pattern |
|---|---|---|---|---|
| mamba2_scan_backward | 75.0% | 1567 | 25 | 1 thread/weight, loop B×K×SH2 |
| q_denoise_backward | 7.3% | 152 | 25 | 1 thread/weight, loop B |
| iqn_bias_grad_reduce | 3.9% | 16 | 130 | 1 thread/neuron, loop B×Q |
| kan_grad_reduce | 3.2% | 8 | 208 | 1 thread/param, loop B |
| curiosity_fwd_bwd_per_block | 2.1% | 1085 | 1 | per-block reduction |
| mamba2_temporal_scan | 1.1% | 22 | 26 | 1 thread/sample (OK), inner matmul SH2 |
| risk_budget_forward | 1.0% | 20 | 26 | 1 thread/sample, inner matmul SH2 |
Design
1. Mamba2 Scan — cuBLAS Projections + Lightweight Scan (75% → <0.5%)
Current: One thread per weight element (12,288 threads), each loops over all 8192 samples. Each iteration does inner products over SH2=256 dimensions. Total: 12,288 × 8192 × 8 × 256 = 206 billion operations on 12,288 threads.
Reference: Official Mamba CUDA kernel (state-spaces/mamba) uses
grid=(batch, dim) — one block per (sample, feature). Uses cub::BlockScan for
sequential dependency. Our kernel uses the exact opposite pattern.
Source: https://github.com/state-spaces/mamba/blob/main/csrc/selective_scan/selective_scan_bwd_kernel.cuh
New (Approach C — cuBLAS projections + lightweight scan):
Forward (mamba2_temporal_scan):
Step 1: Pre-project all timesteps with cuBLAS (3 GEMMs):
A_vals[B×K, STATE_D] = h_history[B×K, SH2] @ W_A[SH2, STATE_D] // m=16, n=B×K, k=256
B_vals[B×K, STATE_D] = h_history[B×K, SH2] @ W_B[SH2, STATE_D]
C_vals[B, STATE_D] = h_s2[B, SH2] @ W_C[SH2, STATE_D]
Step 2: Lightweight scan kernel — grid=(B, ceil(STATE_D/32)), block=(32):
// Each thread handles one (sample, state_dim) pair
// K=8 sequential steps — trivial, no inner matmul
int b = blockIdx.x;
int s = blockIdx.y * blockDim.x + threadIdx.x;
float x = 0.0f;
for (int t = 0; t < K; t++) {
float gate = sigmoid(A_vals[b*K*D + t*D + s]);
x = gate * x + B_vals[b*K*D + t*D + s];
}
// temporal_context = C_vals[b,s] * x
float context_s = C_vals[b*D + s] * x;
// Accumulate into h_enriched via atomicAdd or shared memory reduce
Step 3: Apply temporal context to h_enriched (element-wise kernel or cuBLAS).
Backward (mamba2_scan_backward):
Step 1: Pre-project (same 3 cuBLAS GEMMs as forward — replay activations).
Step 2: Lightweight reverse scan kernel — grid=(B, STATE_D):
// Forward replay + reverse scan for d_x
// Each thread: one (sample, state_dim) pair
// K=8 steps forward, K=8 steps reverse — 16 iterations, no inner matmul
Step 3: Weight gradient accumulation via cuBLAS:
d_W_A[SH2, STATE_D] = h_history^T[SH2, B×K] @ d_gate[B×K, STATE_D] // one GEMM
d_W_B[SH2, STATE_D] = h_history^T[SH2, B×K] @ d_x[B×K, STATE_D] // one GEMM
d_W_C[SH2, STATE_D] = h_s2^T[SH2, B] @ d_context[B, STATE_D] // one GEMM
Total: 6 cuBLAS GEMMs + 2 lightweight scan kernels. Expected: 1567ms + 22ms → ~5ms total.
2. Q-Denoise Backward (7.3% → <0.1%)
Current: 2-layer MLP backward (12→24→12), one thread per weight (1800 params), loops over B=8192.
New: Replace with 2× cuBLAS backward GEMMs + 2× bias_grad_reduce. Same as the main DQN backward pass. Dimensions are tiny (24×24, 12×24) but batch=8192 makes it a proper GEMM.
3. IQN Bias Grad Reduce (3.9% → <0.2%)
Current: 130 calls × 16ms. Each reduces [B×Q, out_dim] over batch. One thread per output neuron, loops over B×Q=524,288. With out_dim=3, that's 3 threads doing 524K iterations each.
New: 2-phase warp reduction:
- Phase 1: grid=(ceil(B×Q/256), out_dim), block=256. Each block sums 256 elements.
- Phase 2: reduce block partials → final result.
Or use
cub::DeviceSegmentedReducefor segmented reduction.
4. KAN Grad Reduce (3.2% → <0.1%)
Same pattern as IQN — one thread per param, loop over batch. Fix: batch-parallel reduction with warp shuffle.
5. Curiosity Fwd+Bwd (2.1% → <0.1%)
Current: 1 instance × 1085ms. Check grid dimensions — likely too few blocks.
New: 2× cuBLAS GEMMs for the curiosity MLP (already cuBLAS for inference, but training kernel is custom). Replace with cuBLAS forward + backward + element-wise loss.
6. Risk Budget Forward (1.0% → <0.01%)
Current: One thread per sample, inner matmul loop over SH2=256.
New: 2× cuBLAS GEMMs. Concat [h_s2; ISV; pred_error] → GEMM1 → ReLU → GEMM2 → sigmoid. Same pattern as the main DQN forward pass.
Scratch Buffer Requirements
Mamba2 pre-projection needs:
A_vals[B×K, STATE_D]= 8192 × 8 × 16 × 4 = 4MBB_vals[B×K, STATE_D]= 4MBC_vals[B, STATE_D]= 0.5MBd_gate[B×K, STATE_D]= 4MB (backward only)d_x[B×K, STATE_D]= 4MB (backward only)
Total: ~17MB additional VRAM. Negligible on H100.
Success Criteria
- Step time: 1950ms → <50ms
- mamba2_scan_backward: 1567ms → <3ms
- All custom kernels use batch-parallel patterns (no per-weight batch loops)
- 899 unit tests + 20 smoke tests pass
- Training loss and grad_norm match within 5% of current values
Files Modified
crates/ml/src/cuda_pipeline/experience_kernels.cu— mamba2_scan_backward, q_denoise_backward, kan_grad_reduce, risk_budget_forwardcrates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu— mamba2_temporal_scancrates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu— iqn_bias_grad_reducecrates/ml/src/cuda_pipeline/curiosity_training_kernel.cu— curiosity_fwd_bwd_per_blockcrates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs— new cuBLAS GEMM descriptors for Mamba2 projections, denoise, risk_budgetcrates/ml/src/trainers/dqn/fused_training.rs— updated capture for new kernel topology