perf(cuda): keep the Mamba2 scan state in registers, stage tokens by chunk #90
Loading…
Reference in a new issue
No description provided.
Delete branch "perf/ssm-scan-chunked"
Deleting a branch is permanent. Although the deleted branch may continue to exist for a short time before it actually gets removed, it CANNOT be undone in most cases. Continue?
A faster Mamba2 scan on CUDA, for prefill and decode alike, on both an RTX PRO 6000 and an RTX 3070. The kernel is restructured but computes the same per-token recurrence, so the numerics and tolerances are unchanged.
Why
The first scan kernel (#88) gave each state element its own thread and reduced
y = h·Cacross the whole block for every token. That meant about 15__syncthreadsper token, and every head-dim row re-read B and C from global memory. A long prefill paid thousands of sequential block-wide reductions per layer.What changes
yis a shuffle reduction inside the warp, so a token needs no block synchronization.__syncthreadsper chunk, and each value is read from global memory once per block rather than once per row.This is still the sequential recurrence the CPU reference defines; only float association changes. It is deliberately not Mamba2's SSD chunked-matmul algorithm, which would also parallelize prefill over time. That rewrite is only worth its extra numerics if the scan is a large share of a real prefill, which this PR does not measure.
Results
One layer's scan at Nemotron 3 Nano's geometry (64 heads × 64 × 128 state, 8 groups), measured by
BenchmarkCUDASSMScan*: prefill over 10 iterations, decode over 2 × 300. The RTX 3070 was selected withCUDA_VISIBLE_DEVICES; its figures vary about ±10% run to run, so they are given as ranges.Across the model's 23 Mamba layers, a 4096-token prefill spends about 80 ms in the scan on the RTX PRO 6000, down from about 238 ms.
How the layout was chosen
The first version of this PR loaded each lane's state slice straight from global memory, with no tile. Adjacent lanes then read 128 bytes apart, and decode got slower than
mainon both cards: 61 → 73 µs on the RTX PRO 6000 and 330 → 429 µs on the 3070. A decode step is almost nothing but that state traffic. It was caught on the 3070 and fixed here before merge.The variants measured:
mainrownsr, r+R, …) avoids the bank conflicts without padding, and decodes about 25% faster. But its long prefills ran about 2× slower on the RTX PRO 6000.mainand the first push everywhere.Testing
TestCUDASSMScanChunksAndRowBlockscovers what the layout adds, against the CPU reference on mixed fresh and resuming batches:go test -tags cuda ./...(24 packages) pass.TestCUDAMatMulNVFP4W8A8skips on the 3070 as designed, since Ampere has no FP8 GEMM.mamba_bench_cuda_test.goholds the benchmarks. The shared test helpersdev/readnow taketesting.TBso benchmarks can use them.go build,go vetandgo test ./...pass.Not in this PR
cudaMallocs its staging and synchronizes on each launch.🤖 Generated with Claude Code
Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id
r-b958e3-f5d083This is an AI-generated review and may contain mistakes.
Status: ❌ Failed
This review couldn't be completed: that model isn't loaded on the inference service right now, and no alternate model was able to review it either. Consider splitting this PR into smaller changes. Tracking id
r-b958e3-f5d083.Comment
@pr-reviewer-bot retryto try again.11d17175d5952514c1b0