perf(cuda): keep the Mamba2 scan state in registers, stage tokens by chunk #90

Merged
rcsheets merged 1 commit from perf/ssm-scan-chunked into main 2026-09-27 18:11:09 +00:00
Owner

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·C across the whole block for every token. That meant about 15 __syncthreads per 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

  • State lives in registers. Each lane holds 32 contiguous state elements. A row is split across the adjacent lanes of one warp: the power of two that covers the state size, which is 4 lanes at Nemotron's 128. y is a shuffle reduction inside the warp, so a token needs no block synchronization.
  • Tokens are staged by chunk. A block covers one (sequence, head, run of rows). It stages a chunk of tokens' B, C, dt and its rows' x into shared memory: 25 tokens at Nemotron's size, within 32 KiB. That is two __syncthreads per chunk, and each value is read from global memory once per block rather than once per row.
  • B and C are padded. They are staged with 4 floats of padding per 32, so each lane's contiguous slice starts on its own shared-memory bank. Unpadded, a row's lanes read 32 floats apart, all on one bank, and serialize.
  • State moves through a coalesced tile. A block's rows are contiguous in the state slot, so the state moves between the slot and the registers through a shared-memory tile, with one coalesced copy each way. The tile reuses the staging buffer.

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 with CUDA_VISIBLE_DEVICES; its figures vary about ±10% run to run, so they are given as ranges.

RTX PRO 6000 (sm_120, 300 W cap) RTX 3070 (sm_86)
prefill, 1 × 4096 tokens 10.3 → 3.5 ms (2.9×) ~52 → 6.0–6.4 ms (~8×)
decode, 16 × 1 token 61 → 32 µs (1.9×) 330 → 147–163 µs (~2×)

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 main on 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:

variant PRO 6000 prefill / decode 3070 prefill / decode
main 10.3 ms / 61 µs ~52 ms / 330 µs
contiguous slices, no tile (first push) 5.1 ms / 73 µs 10.1 ms / 429 µs
interleaved elements + tile 7.7 ms / 25 µs 5.9 ms / 113 µs
contiguous + padding + tile (this PR) 3.5 ms / 32 µs 6.0–6.4 ms / 147–163 µs
  • Interleaving each row's elements across its lanes (lane r owns r, 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.
  • Padded contiguous slices are the only variant faster than main and the first push everywhere.
  • Making the lane count a compile-time constant measured identically and was dropped.

Testing

  • New test. TestCUDASSMScanChunksAndRowBlocks covers what the layout adds, against the CPU reference on mixed fresh and resuming batches:
    • sequences crossing several staging chunks (up to 150 tokens)
    • a head whose rows span two blocks (96 rows)
    • a state size that leaves a lane's slice partly empty (100)
    • 32 lanes per row (state size 1024)
  • On both cards: all seven Mamba op tests, the five Nemotron-H GPU tests (goldens, NVFP4, BF16) and the full go test -tags cuda ./... (24 packages) pass. TestCUDAMatMulNVFP4W8A8 skips on the 3070 as designed, since Ampere has no FP8 GEMM.
  • Benchmarks. mamba_bench_cuda_test.go holds the benchmarks. The shared test helpers dev/read now take testing.TB so benchmarks can use them.
  • Pure-Go build: go build, go vet and go test ./... pass.

Not in this PR

  • Decode's remaining state cost. The tile-to-register copy still bank-conflicts in the contiguous layout. A swizzled tile would remove that, but it exceeds the 48 KiB shared-memory default at some state sizes.
  • Measuring the scan's share of a real prefill. This needs serving the real checkpoint with a long prompt, before and after, and it decides whether the SSD rewrite is worth doing.
  • Per-launch overhead. Every op cudaMallocs its staging and synchronizes on each launch.
  • Routed experts. Each picked expert still runs over every token in the step.

🤖 Generated with Claude Code

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·C` across the whole block for every token. That meant about 15 `__syncthreads` per 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 - **State lives in registers.** Each lane holds 32 contiguous state elements. A row is split across the adjacent lanes of one warp: the power of two that covers the state size, which is 4 lanes at Nemotron's 128. `y` is a shuffle reduction inside the warp, so a token needs no block synchronization. - **Tokens are staged by chunk.** A block covers one (sequence, head, run of rows). It stages a chunk of tokens' B, C, dt and its rows' x into shared memory: 25 tokens at Nemotron's size, within 32 KiB. That is two `__syncthreads` per chunk, and each value is read from global memory once per block rather than once per row. - **B and C are padded.** They are staged with 4 floats of padding per 32, so each lane's contiguous slice starts on its own shared-memory bank. Unpadded, a row's lanes read 32 floats apart, all on one bank, and serialize. - **State moves through a coalesced tile.** A block's rows are contiguous in the state slot, so the state moves between the slot and the registers through a shared-memory tile, with one coalesced copy each way. The tile reuses the staging buffer. 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 with `CUDA_VISIBLE_DEVICES`; its figures vary about ±10% run to run, so they are given as ranges. | | RTX PRO 6000 (sm_120, 300 W cap) | RTX 3070 (sm_86) | |---|---|---| | prefill, 1 × 4096 tokens | 10.3 → **3.5 ms** (2.9×) | ~52 → **6.0–6.4 ms** (~8×) | | decode, 16 × 1 token | 61 → **32 µs** (1.9×) | 330 → **147–163 µs** (~2×) | 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 `main` on 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: | variant | PRO 6000 prefill / decode | 3070 prefill / decode | |---|---|---| | `main` | 10.3 ms / 61 µs | ~52 ms / 330 µs | | contiguous slices, no tile (first push) | 5.1 ms / 73 µs | 10.1 ms / 429 µs | | interleaved elements + tile | 7.7 ms / 25 µs | 5.9 ms / 113 µs | | **contiguous + padding + tile (this PR)** | **3.5 ms / 32 µs** | **6.0–6.4 ms / 147–163 µs** | - **Interleaving** each row's elements across its lanes (lane `r` owns `r, 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. - **Padded contiguous slices** are the only variant faster than `main` and the first push everywhere. - **Making the lane count a compile-time constant** measured identically and was dropped. ## Testing - **New test.** `TestCUDASSMScanChunksAndRowBlocks` covers what the layout adds, against the CPU reference on mixed fresh and resuming batches: - sequences crossing several staging chunks (up to 150 tokens) - a head whose rows span two blocks (96 rows) - a state size that leaves a lane's slice partly empty (100) - 32 lanes per row (state size 1024) - **On both cards:** all seven Mamba op tests, the five Nemotron-H GPU tests (goldens, NVFP4, BF16) and the full `go test -tags cuda ./...` (24 packages) pass. `TestCUDAMatMulNVFP4W8A8` skips on the 3070 as designed, since Ampere has no FP8 GEMM. - **Benchmarks.** `mamba_bench_cuda_test.go` holds the benchmarks. The shared test helpers `dev`/`read` now take `testing.TB` so benchmarks can use them. - **Pure-Go build:** `go build`, `go vet` and `go test ./...` pass. ## Not in this PR - **Decode's remaining state cost.** The tile-to-register copy still bank-conflicts in the contiguous layout. A swizzled tile would remove that, but it exceeds the 48 KiB shared-memory default at some state sizes. - **Measuring the scan's share of a real prefill.** This needs serving the real checkpoint with a long prompt, before and after, and it decides whether the SSD rewrite is worth doing. - **Per-launch overhead.** Every op `cudaMalloc`s its staging and synchronizes on each launch. - **Routed experts.** Each picked expert still runs over every token in the step. 🤖 Generated with [Claude Code](https://claude.com/claude-code)
perf(cuda): keep the Mamba2 scan state in registers, stage tokens by chunk
All checks were successful
ci / test_and_build (pull_request) Successful in 24s
11d17175d5
The first scan kernel gave each state element its own thread and
reduced y = h.C across the block for every token: about 15
__syncthreads per token, with B and C re-read from global memory by
every head-dim row. For a long prefill that is thousands of sequential
block-wide reductions per layer.

The scan now keeps each row's state in registers, 32 elements per lane,
and splits a row across the adjacent lanes of one warp (the power of
two that covers the state size). y is a shuffle reduction inside the
warp, so a token needs no block synchronization. A block covers one
(sequence, head, run of rows) and stages a chunk of tokens' B, C, dt and
the rows' x into shared memory, up to 64 tokens in 32 KiB (25 at
Nemotron's size). That takes two __syncthreads per chunk, and each
value is read from global memory once per block rather than once per
row.

It is the same per-token recurrence the CPU reference defines. Only
float association changes, so the existing tolerances hold.

At Nemotron 3 Nano's geometry (64 heads x 64 x 128 state, 8 groups) on
the RTX PRO 6000, BenchmarkCUDASSMScanPrefill4096 (one layer, 4096
tokens) goes from 10.3 ms to 5.1 ms, about 238 ms to 117 ms across the
23 Mamba layers. Decode (16 sequences x 1 token) is unchanged at about
87 us, where the launch's staging copy and synchronization, not the
scan, set the cost.

TestCUDASSMScanChunksAndRowBlocks covers what the layout adds:
sequences crossing several staging chunks, a head whose rows span two
blocks, a state size that leaves a lane's slice partly empty, and 32
lanes per row (state size 1024). The Mamba op tests, the Nemotron-H GPU
tests and the full CUDA suite pass. The benchmark helpers take
testing.TB so benchmarks can share them.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Collaborator

Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id r-b958e3-f5d083
This 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 retry to try again.

<!-- pr-reviewer:review --> *Automated review by [pr-reviewer](https://git.brooktrails.org/brooktrails/pr-reviewer) v0.52.3 | Safety Check | Mistral Small | tracking id `r-b958e3-f5d083`* *This 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 retry` to try again.
rcsheets force-pushed perf/ssm-scan-chunked from 11d17175d5
All checks were successful
ci / test_and_build (pull_request) Successful in 24s
to 952514c1b0
All checks were successful
ci / test_and_build (pull_request) Successful in 26s
2026-09-27 17:56:50 +00:00
Compare
rcsheets deleted branch perf/ssm-scan-chunked 2026-09-27 18:11:10 +00:00
Sign in to join this conversation.
No reviewers
No labels
No milestone
No project
No assignees
2 participants
Notifications
Due date
The due date is invalid or out of range. Please use the format "yyyy-mm-dd".

No due date set.

Dependencies

No dependencies set

Reference
brooktrails/gllm!90
No description provided.