feat(cuda): Mamba2 kernels for hybrid models #88
Loading…
Reference in a new issue
No description provided.
Delete branch "feat/mamba2-cuda"
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?
Phase 5 of
docs/hybrid-state-cache.md: CUDA kernels for the four Mamba2 ops. The ops had CPU references (#85) but returnedErrNotImplementedon 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
relu2causal_conv1dkernel-1inputs (zeros when fresh) and writes the new ones back. The kernel width is capped at 8; Mamba2 uses 4ssm_scanhin a register across the sequence's tokens, with a block reduction per token fory = h·Cgated_rmsnormssm_scanis 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)
make test-cudapasses in full (about 8 s).go build,go vetandgo test ./...pass on the pure-Go build.The real checkpoint
gllm serveran NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4 on the RTX PRO 6000 with--kv-cache 8GiB --max-model-len 4096:reasoning_content("Straightforward: Paris. Provide one sentence answer…"), then the answer "The capital of France is Paris." incontent, finishing on<|im_end|>.enable_thinking: false. It answered a second question directly, with no reasoning field.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
mistralloader acts onbackend.BF16Weights;mistral4and 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.BF16Residentmarks an architecture whose loader keeps BF16 when asked.applyWeightDTypeand the planner check it:autokeeps BF16 only for such an architecture. Otherwise it resolves to f32, with no misleading log line and no scratch charge.--weights-dtype bf16for an architecture that widens is refused, naming the architecture, instead of being a silent no-op.What the
nemotronhloader keeps. It now implements the interface and keeps BF16 for the matrices MatMul and Embedding read in BF16:lm_headin_projrow blocks andout_projIt 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-nvfp4now 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:
TestBF16WeightsAreResidentpins which tensors stay BF16 and checks that the BF16 load is smaller.TestForwardNVFP4BF16CUDArequires the same argmax and at most 2% relative L2 error despite the bf16 GEMM's activation rounding, as mistral's BF16 GPU test does.auto→ f32,bf16refused).make test-cudapasses.Not in this PR
🤖 Generated with Claude Code
Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id
r-b913dd-fc8d3fThis 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 retryto try again.