feat(cuda): Mamba2 kernels for hybrid models #88

Merged
rcsheets merged 3 commits from feat/mamba2-cuda into main 2026-09-27 13:06:10 +00:00
Owner

Phase 5 of docs/hybrid-state-cache.md: CUDA kernels for the four Mamba2 ops. The ops had CPU references (#85) but returned ErrNotImplemented on CUDA, so a hybrid model on the GPU failed its first forward with a 501. With this PR, Nemotron-H runs on the GPU and matches the same HF goldens as the CPU.

Kernels

op parallelism
relu2 elementwise
causal_conv1d one thread per (sequence, channel), walking the sequence's rows with the conv window in registers. It seeds from the slot's last kernel-1 inputs (zeros when fresh) and writes the new ones back. The kernel width is capped at 8; Mamba2 uses 4
ssm_scan one block per (sequence, head, head-dim row), one thread per state element, holding its element of h in a register across the sequence's tokens, with a block reduction per token for y = h·C
gated_rmsnorm one block per (row, group)

ssm_scan is the per-token recurrence that the CPU reference defines. It is correctness-first: a long prefill walks its tokens in order. The chunked prefill scan is later throughput work, and will be validated against this kernel.

Each sequence's first row, row count, slot and fresh flag are passed as one packed host array, staged to the device once per launch as the other ops stage their index arrays.

The Go wrappers repeat the CPU reference's checks: slot in range, no slot shared by two sequences, run lengths covering the rows. On the GPU a bad batch is therefore refused before it can become an out-of-bounds write or two sequences racing on one state.

Validation (RTX PRO 6000, sm_120)

  • Each op against the CPU reference. Batches are ragged and mix a fresh sequence (its slot holding garbage it must ignore) with sequences resuming from random state, on slots in a different order from their rows. Outputs and every slot are compared.
  • Edge shapes:
    • the scan at state sizes 16 (narrower than a warp), 128 (Nemotron's) and 20 (not a power of two)
    • the conv over 300 channels (more than one block per sequence), with a one-token run shorter than its history
  • Split versus whole: the scan run in two pieces matches the one-step run on the GPU.
  • Rejection: bad batches are rejected on the GPU as on the CPU.
  • End to end: the Nemotron-H tiny fixtures run on CUDA against the same HF goldens as the CPU, within 2e-4, covering prefill, token-by-token decode, two sequences batched on separate slots, and the modelopt NVFP4 fixture.

make test-cuda passes in full (about 8 s). go build, go vet and go test ./... pass on the pure-Go build.

The real checkpoint

gllm serve ran NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4 on the RTX PRO 6000 with --kv-cache 8GiB --max-model-len 4096:

  • Load. It loaded in 10 s: 20.0 GiB of weights, 0.75 GiB of KV, and 16 × 47.6 MiB of recurrent state.
  • Thinking left to the template's default (on). A greedy request for France's capital produced 35 reasoning tokens in reasoning_content ("Straightforward: Paris. Provide one sentence answer…"), then the answer "The capital of France is Paris." in content, finishing on <|im_end|>.
  • enable_thinking: false. It answered a second question directly, with no reasoning field.
  • Speed. Decode ran at about 42 ms per token under the card's 300 W cap.

The second commit records this in the README and the design doc.

BF16 residency (third commit)

The serve run logged "keeping BF16 checkpoint weights unwidened", but it held 20.0 GiB of weights widened to F32. Only the mistral loader acts on backend.BF16Weights; mistral4 and GLM have the same gap. The forward scratch reserve was also charged for a bf16 activation buffer no MatMul used.

  • Which architectures keep BF16. model.BF16Resident marks an architecture whose loader keeps BF16 when asked. applyWeightDType and the planner check it:

    • auto keeps BF16 only for such an architecture. Otherwise it resolves to f32, with no misleading log line and no scratch charge.
    • An explicit --weights-dtype bf16 for an architecture that widens is refused, naming the architecture, instead of being a silent no-op.
  • What the nemotronh loader keeps. It now implements the interface and keeps BF16 for the matrices MatMul and Embedding read in BF16:

    • the embedding and lm_head
    • the attention projections
    • unquantized in_proj row blocks and out_proj

    It widens everything that feeds an F32-only op: the norms, the conv weight and the per-head Mamba parameters. The router gate stays F32, because the reference scores experts in float32.

  • Result: the real checkpoint's weights drop from 20.0 to 18.0 GiB (gllm plan), which is what the log already said.

  • Fixture. tiny-nvfp4 now stores its unquantized tensors as BF16, as the real checkpoint does, except the router gate and bias. Its goldens are recomputed from the BF16-rounded weights.

  • Tests:

    • On the CPU, the fixture matches its goldens both widened and BF16-resident, at the same tolerance.
    • TestBF16WeightsAreResident pins which tensors stay BF16 and checks that the BF16 load is smaller.
    • On the GPU, TestForwardNVFP4BF16CUDA requires the same argmax and at most 2% relative L2 error despite the bf16 GEMM's activation rounding, as mistral's BF16 GPU test does.
    • The engine and planner tests cover the widening-architecture cases (auto → f32, bf16 refused).
    • make test-cuda passes.

Not in this PR

  • Throughput: the chunked prefill scan, and gathered routed experts instead of per-expert passes.

🤖 Generated with Claude Code

Phase 5 of `docs/hybrid-state-cache.md`: CUDA kernels for the four Mamba2 ops. The ops had CPU references (#85) but returned `ErrNotImplemented` on CUDA, so a hybrid model on the GPU failed its first forward with a 501. With this PR, Nemotron-H runs on the GPU and matches the same HF goldens as the CPU. ## Kernels | op | parallelism | |---|---| | `relu2` | elementwise | | `causal_conv1d` | one thread per (sequence, channel), walking the sequence's rows with the conv window in registers. It seeds from the slot's last `kernel-1` inputs (zeros when fresh) and writes the new ones back. The kernel width is capped at 8; Mamba2 uses 4 | | `ssm_scan` | one block per (sequence, head, head-dim row), one thread per state element, holding its element of `h` in a register across the sequence's tokens, with a block reduction per token for `y = h·C` | | `gated_rmsnorm` | one block per (row, group) | `ssm_scan` is the per-token recurrence that the CPU reference defines. It is correctness-first: a long prefill walks its tokens in order. The chunked prefill scan is later throughput work, and will be validated against this kernel. Each sequence's first row, row count, slot and fresh flag are passed as one packed host array, staged to the device once per launch as the other ops stage their index arrays. The Go wrappers repeat the CPU reference's checks: slot in range, no slot shared by two sequences, run lengths covering the rows. On the GPU a bad batch is therefore refused before it can become an out-of-bounds write or two sequences racing on one state. ## Validation (RTX PRO 6000, sm_120) - **Each op against the CPU reference.** Batches are ragged and mix a fresh sequence (its slot holding garbage it must ignore) with sequences resuming from random state, on slots in a different order from their rows. Outputs and every slot are compared. - **Edge shapes:** - the scan at state sizes 16 (narrower than a warp), 128 (Nemotron's) and 20 (not a power of two) - the conv over 300 channels (more than one block per sequence), with a one-token run shorter than its history - **Split versus whole:** the scan run in two pieces matches the one-step run on the GPU. - **Rejection:** bad batches are rejected on the GPU as on the CPU. - **End to end:** the Nemotron-H tiny fixtures run on CUDA against the same HF goldens as the CPU, within 2e-4, covering prefill, token-by-token decode, two sequences batched on separate slots, and the modelopt NVFP4 fixture. `make test-cuda` passes in full (about 8 s). `go build`, `go vet` and `go test ./...` pass on the pure-Go build. ## The real checkpoint `gllm serve` ran NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4 on the RTX PRO 6000 with `--kv-cache 8GiB --max-model-len 4096`: - **Load.** It loaded in 10 s: 20.0 GiB of weights, 0.75 GiB of KV, and 16 × 47.6 MiB of recurrent state. - **Thinking left to the template's default (on).** A greedy request for France's capital produced 35 reasoning tokens in `reasoning_content` ("Straightforward: Paris. Provide one sentence answer…"), then the answer "The capital of France is Paris." in `content`, finishing on `<|im_end|>`. - **`enable_thinking: false`.** It answered a second question directly, with no reasoning field. - **Speed.** Decode ran at about 42 ms per token under the card's 300 W cap. The second commit records this in the README and the design doc. ## BF16 residency (third commit) The serve run logged "keeping BF16 checkpoint weights unwidened", but it held 20.0 GiB of weights widened to F32. Only the `mistral` loader acts on `backend.BF16Weights`; `mistral4` and GLM have the same gap. The forward scratch reserve was also charged for a bf16 activation buffer no MatMul used. - **Which architectures keep BF16.** `model.BF16Resident` marks an architecture whose loader keeps BF16 when asked. `applyWeightDType` and the planner check it: - `auto` keeps BF16 only for such an architecture. Otherwise it resolves to f32, with no misleading log line and no scratch charge. - An explicit `--weights-dtype bf16` for an architecture that widens is refused, naming the architecture, instead of being a silent no-op. - **What the `nemotronh` loader keeps.** It now implements the interface and keeps BF16 for the matrices MatMul and Embedding read in BF16: - the embedding and `lm_head` - the attention projections - unquantized `in_proj` row blocks and `out_proj` It widens everything that feeds an F32-only op: the norms, the conv weight and the per-head Mamba parameters. The router gate stays F32, because the reference scores experts in float32. - **Result:** the real checkpoint's weights drop from 20.0 to 18.0 GiB (`gllm plan`), which is what the log already said. - **Fixture.** `tiny-nvfp4` now stores its unquantized tensors as BF16, as the real checkpoint does, except the router gate and bias. Its goldens are recomputed from the BF16-rounded weights. - **Tests:** - On the CPU, the fixture matches its goldens both widened and BF16-resident, at the same tolerance. - `TestBF16WeightsAreResident` pins which tensors stay BF16 and checks that the BF16 load is smaller. - On the GPU, `TestForwardNVFP4BF16CUDA` requires the same argmax and at most 2% relative L2 error despite the bf16 GEMM's activation rounding, as mistral's BF16 GPU test does. - The engine and planner tests cover the widening-architecture cases (`auto` → f32, `bf16` refused). - `make test-cuda` passes. ## Not in this PR - **Throughput:** the chunked prefill scan, and gathered routed experts instead of per-expert passes. 🤖 Generated with [Claude Code](https://claude.com/claude-code)
feat(cuda): Mamba2 kernels for hybrid models
All checks were successful
ci / test_and_build (pull_request) Successful in 25s
7796098f9f
Phase 5 of docs/hybrid-state-cache.md. The four Mamba2 ops had CPU
references (internal/backend/cpu/mamba.go) and returned ErrNotImplemented
on CUDA. They now have kernels:

- relu2: elementwise.
- causal_conv1d: one thread per (sequence, channel), walking the
  sequence's rows with the conv window in registers. It seeds from the
  slot's last kernel-1 inputs (zeros when fresh) and writes the new ones
  back. The kernel width is capped at 8; Mamba2 uses 4.
- ssm_scan: one block per (sequence, head, head-dim row), one thread
  per state element, holding its element of h in a register across the
  sequence's tokens. Each token does a block reduction for y = h.C. This
  is the per-token recurrence the CPU reference defines; a chunked
  prefill scan is later throughput work, validated against this one.
- gated_rmsnorm: one block per (row, group).

Each sequence's first row, row count, slot and fresh flag go to the
kernels as one packed host array, staged to the device once per launch
like the other ops' index arrays. The Go wrappers repeat the CPU
reference's checks (slot in range, no slot shared by two sequences, run
lengths covering the rows), so on the GPU a bad batch is refused before
it can become an out-of-bounds write or a race between two sequences.

Validated on the RTX PRO 6000 (sm_120):

- Each op against the CPU reference, on ragged batches that mix a fresh
  sequence (its slot holding garbage it must ignore) with sequences
  resuming from random state, on slots in a different order from their
  rows. Outputs and every slot are compared.
- The scan at state sizes 16 (narrower than a warp), 128 (Nemotron's)
  and 20 (not a power of two). The conv over 300 channels (more than
  one block) with a one-token run shorter than its history.
- A split-versus-whole scan on the GPU.
- The Nemotron-H tiny fixtures on CUDA against the same HF goldens as
  the CPU, within 2e-4: prefill, token-by-token decode, two sequences
  batched on separate slots, and the modelopt NVFP4 fixture.

make test-cuda passes in full.

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-b913dd-fc8d3f
This is an AI-generated review and may contain mistkaes.

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-b913dd-fc8d3f.

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-b913dd-fc8d3f`* *This is an AI-generated review and may contain mistkaes.* **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-b913dd-fc8d3f`. Comment `@pr-reviewer-bot retry` to try again.
docs: record the real Nemotron 3 checkpoint serving on the GPU
All checks were successful
ci / test_and_build (pull_request) Successful in 25s
f730bec5b1
Loaded in 10 s on one RTX PRO 6000 (20.0 GiB of weights, 0.75 GiB of KV,
16 x 47.6 MiB of recurrent state). With thinking left to the template's
default, a greedy chat turn reasoned (35 tokens, in reasoning_content) and
answered "The capital of France is Paris.", ending on <|im_end|>. With
enable_thinking false it answered directly. Decode ran at about 42 ms per
token under the card's 300 W cap.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
fix(engine): keep BF16 resident only where the loader honors it
All checks were successful
ci / test_and_build (pull_request) Successful in 26s
96c0c45352
The engine logged "keeping BF16 checkpoint weights unwidened" for any
BF16 checkpoint on a capable backend, but only the mistral loader acts
on backend.BF16Weights. Serving Nemotron 3 Nano logged the 18.0 GiB
footprint while holding 20.0 GiB of weights widened to F32. mistral4 and
GLM have the same gap. The forward scratch reserve was also charged for
a bf16 activation buffer no MatMul would use.

model.BF16Resident marks an architecture whose loader keeps BF16 when
asked. applyWeightDType and the planner now check it:

- auto keeps BF16 only for such an architecture. Otherwise it resolves
  to f32, without the log line or the scratch charge.
- An explicit --weights-dtype bf16 for an architecture that widens is
  refused and names the architecture, where it used to be a silent
  no-op.

mistral and nemotron_h implement it. The nemotronh loader now keeps the
BF16 matrices that MatMul and Embedding read in BF16: the embedding,
lm_head, the attention projections, and unquantized in_proj row blocks
and out_proj. It widens everything that feeds an F32-only op (norms, the
conv weight, the per-head Mamba parameters), and the router gate stays
F32, since the reference scores experts in float32. The real
checkpoint's resident weights drop from 20.0 to 18.0 GiB (gllm plan),
matching what the log already said.

The tiny-nvfp4 fixture now stores its unquantized tensors as BF16,
except the router gate and bias, as the real checkpoint does. Its
goldens are recomputed from the BF16-rounded weights. On the CPU it
matches them widened and BF16-resident alike, at the same tolerance.
TestBF16WeightsAreResident pins which tensors stay BF16 and that the
load is smaller. On the GPU, TestForwardNVFP4BF16CUDA holds argmax and
a 2% relative L2 under the bf16 GEMM's activation rounding, the same
check as mistral's BF16 GPU test. The engine and planner tests cover
the widening-architecture cases.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
rcsheets deleted branch feat/mamba2-cuda 2026-09-27 13:06: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!88
No description provided.