Repository navigation
[SM90] Add native block32 FP8 GEMM support (dsv41) - #95
Conversation
Support recipe=(1,32,32) in dense NT FP8 GEMM using four independently scaled K32 partial sums per K128 TMA stage. Adapt scale layouts, accumulator scheduling, tile heuristics, and N/K specialization while retaining the default block128 path.
Move block32 capability checks from the shared scale-layout helper to GEMM dispatch. Derive scale counts and offsets from block parameters, centralize accumulator-bank selection, clarify N/K specialization, and fix block128 branch indentation. Validated on H200: 65 existing gates and 8 API checks passed; all 360 dense block32 points match the previous commit bit for bit. Paired full-path timings across 40 points give a 1.003x geometric mean baseline/cleanup ratio.
|
Thanks for adding native block32 support! I reviewed head Please also validate the activation M granularity in the SM90 dispatch. The new layout branch accepts Could you also commit regression tests for the new SM90 block32 path? The current Hopper generator only produces the legacy FP8 block128 configuration. In particular, please cover independent dequantized-reference correctness, partial K=128 stages, N/M tails, both B-scale layouts, and all three accumulator-bank schedules. The existing block128 tests do not exercise the new implementation. |
Require one activation scale per row in the SM90 dense FP8 dispatch for both block32 and block128. Add Hopper regression coverage with an independent dequantized reference, K-stage fragments, M/N tails, both B-scale layouts, all three accumulator-bank schedules, and invalid recipe rejection.
Run native block32 correctness, tail, scale-layout, bank-schedule, and invalid activation-scale cases from the existing FP8 GEMM test entrypoint. Remove the dedicated SM90 test file and runner dispatch; keep the existing block128 positive coverage in the general suite.
|
@BBuf btw it looks like your find with (32,32,32) also works for the current release with (128,128,128): import torch
import sgl_deep_gemm as dg
m = n = 32
k = 128
a = torch.ones((m, k), device="cuda").to(torch.float8_e4m3fn)
b = torch.ones((n, k), device="cuda").to(torch.float8_e4m3fn)
sb = torch.ones((1, 1), device="cuda")
def run(sa):
out = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
dg.fp8_gemm_nt((a, sa), (b, sb), out, recipe=(128, 128, 128))
torch.cuda.synchronize()
return out
backing = torch.ones((m, 1), device="cuda")
sa = backing[:1]
before = run(sa)
backing[1:] = 3
after = run(sa)
print(before[0, 0], before[1, 0]) # 128, 128
print(after[0, 0], after[1, 0]) # 128, 384in this case, it is logically assumed that there is exactly one a-scale, but changing values outside of it causes changing the results I added an additional check, which will now change the behavior for block128 too |
Make sense, approved. |
Motivation
Dense Fp8 weights in DeepSeek-V4.1-Flash use
[32,32]blocks, while DeepSeek-V4-Flash used[128,128].The main sources of these dense GEMMs in V4.1 are: attn, shared expert, indexer, engram, dspark linears.
On Hopper, the sglang dispatcher sends K-block != 128 to Triton:
https://github.com/sgl-project/sglang/blob/98fce73d5bd0a25afe7d68443d314190b1c47e64/python/sglang/srt/layers/quantization/fp8_utils.py#L582-L605
and the DeepGEMM wrapper only accepts
[128,128]:https://github.com/sgl-project/sglang/blob/98fce73d5bd0a25afe7d68443d314190b1c47e64/python/sglang/srt/layers/quantization/fp8_utils.py#L1223-L1242
PR sgl-project/sglang#41251 illustrates tuning this Triton path with H200 group32 configs.
This change enables DeepGEMM to consume the original
[32,32]fp8 weights and scales.Implementation
Adds a dense NT FP8 --> BF16 kernel for SM90:
The implementation extends the existing SM90 block128 kernel, reusing WGMMA, TMA, the persistent scheduler, pipeline, and BF16 epilogue. Each TMA stage still loads K=128, but computes four independent K=32 partial sums. Each partial is multiplied by its own activation and weight scales before being added to the final FP32 accumulator.
Benchmarks
H200
The table shows median GEMM latency in µs.
Comparison with tuned Triton32
We compare the end-to-end GPU compute path intended for a dense FP8 linear operation in SGLang: BF16 activations --> block32 FP8 quantization --> GEMM --> BF16 output. Both paths use the same input activations and prequantized weights. Activation quantization and scale generation in the layout required by each backend are included in every timed invocation.
Triton uses the tuned GEMM configuration from PR sgl-project/sglang#41251 and a quantization producer at each point.
By M range: