perf: gather tokens per MoE expert, run large NVFP4 GEMMs on bf16 tensor cores #91
Loading…
Reference in a new issue
No description provided.
Delete branch "perf/moe-gather"
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 4314-token Nemotron 3 Nano prefill on the RTX PRO 6000 goes from 9.97 s to 2.39 s, with its GPU time falling from 8.5 s to 1.0 s. The output is unchanged. The changes came out of profiling a real long prefill of the checkpoint with Nsight Systems.
What the profile showed
Per 4314-token prefill, before this PR (
mainafter #90):cutlass_80_simt_sgemm), ~5,970 launchesTwo problems stacked in that 84%:
nvfp4GEMVMaxRowsrows, the weight was widened to f32 for an f32 GEMM, which cuBLASLt serves with a SIMT kernel.The scan's 1% settles the open question from #90: Mamba2's SSD rewrite is not worth doing.
What changes
MoE gather (
nemotronh). Each (token, picked expert) pair gets a row in packed buffers ofn × topKrows, grouped by expert.Embeddinggathers the inputs.WeightedRowSumadds each token's topK outputs onto the shared expert's.This replaces a ScaleRows and an Add per expert. Two new
backend.Opssupport it, with CPU, CUDA and sizing implementations:RowSlice(t, r0, r1)views a contiguous row range without copying, so experts share buffers allocated once per step.WeightedRowSum(dst, src, rows, weights)computesdst[i] += Σⱼ w[i,j] · src[rows[i,j]]in a fixedjorder: one thread per output element, deterministic, no atomics.forwardScratchgains the packed buffers' term for a pattern's MoE layers. On Nemotron 3 Nano that is the widest layer: 57,984 values per token.bf16 NVFP4 GEMM (CUDA). Above the GEMV threshold,
matMulNVFP4now unpacks to bf16 and runs the bf16 tensor-core GEMM.gllm_dequant_nvfp4_bf16). An e2m1 code times an e4m3 block scale has at most 6 significant bits, so the unpacked weight is exact in bf16.alpha(gllm_lt_matmul_bf16now takes one), where a bf16 weight could only approximate it.Results
Per 4314-token prefill on the RTX PRO 6000, measured with the same Nsight Systems setup before and after:
nvjet_*)cudaMalloc/cudaFreepairs for staging and scratch. After that come prefill attention (0.32 s, the decode flash kernel on a 4k prompt) and the NVFP4 unpacks (0.21 s). The design doc records all of this.Testing
RowSlice(it is a view, writes reach the parent, bounds are checked) andWeightedRowSum(hand values, bounds).WeightedRowSumagainst the CPU, with repeated source rows and a pre-filled dst.RowSlice: a MatMul into a slice lands in exactly those rows of the parent.TestCUDAMatMulNVFP4Exactfeeds bf16-exact activations through the bf16 path and holds it to the f32 tolerance, so a lossy unpack or a missing or doubled global fails.TestCUDAMatMulBF16does.go test -tags cuda ./...passes on the RTX PRO 6000 and the RTX 3070 (24 packages each; FP8 W8A8 skips on the 3070 as designed).go build,go vetandgo test ./...pass.🤖 Generated with Claude Code
Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id
r-b960a7-61ef1eThis 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-b960a7-61ef1e.Comment
@pr-reviewer-bot retryto try again.