Skip to content

[SM90] Add native block32 FP8 GEMM support (dsv41) - #95

Merged
BBuf merged 4 commits into
sgl-project:devfrom
vladnosiv:native-block32-sm90
Oct 1, 2026
Merged

BBuf merged 4 commits into
sgl-project:devfrom
vladnosiv:native-block32-sm90

Conversation

@vladnosiv

Copy link
Copy Markdown

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:

fp8_gemm_nt((Aq, As), (Wq, Ws), D, recipe=(1, 32, 32))

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.

N × K M=1 M=32 M=128 M=512 M=8192
576 × 5120 8.68 8.75 8.89 12.51 65.85
1152 × 5120 9.62 9.06 8.91 17.61 136.06
1792 × 5120 9.24 9.28 9.08 24.98 191.53
4096 × 1280 5.69 4.08 4.60 9.50 120.61
5120 × 288 3.04 3.00 3.20 7.01 54.99
5120 × 576 3.62 3.52 3.75 8.69 81.39
5120 × 1024 5.27 5.18 5.74 11.60 120.76
5120 × 2048 7.93 7.92 8.83 19.37 226.82
8192 × 1280 6.90 5.66 7.83 16.86 238.55
25600 × 6144 59.13 61.98 70.08 231.43 3116.58

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.

N × K Full-path speedup
576 × 5120 1.391×
1152 × 5120 1.594×
1792 × 5120 1.631×
4096 × 1280 1.775×
5120 × 288 1.294×
5120 × 576 1.553×
5120 × 1024 1.787×
5120 × 2048 1.899×
8192 × 1280 1.808×
25600 × 6144 2.025×

By M range:

M Points Full-path speedup
1-4 40 0.951×
8-64 110 1.248×
65-512 100 1.982×
513-8192 110 2.308×

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.
@BBuf
BBuf requested a review from Fridge003 September 30, 2026 08:05
@BBuf

BBuf commented Sep 30, 2026

Copy link
Copy Markdown
Collaborator

Thanks for adding native block32 support! I reviewed head e462f4d. The per-K32 scaling approach looks reasonable. I have two follow-ups before merging:

Please also validate the activation M granularity in the SM90 dispatch. The new layout branch accepts (32, 32) scales for A as well as B, so recipe=(32, 32, 32) passes these checks. However, make_tma_sf_desc still describes A scales using the full M dimension, and the kernel reads one scale per row. For M=N=32, K=128, a compact [1, 4] A-scale tensor is therefore interpreted as [32, 4], reading beyond the logical scale tensor and potentially its allocation. Please require gran_m == 1 for this path and add rejection tests for both recipe forms. This finding comes from source tracing; I have prepared a reproducer but have not executed it on a Hopper GPU.

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.
@vladnosiv

vladnosiv commented Sep 30, 2026 •

Copy link
Copy Markdown
Author

@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, 384

in 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

@BBuf

BBuf commented Sep 30, 2026

Copy link
Copy Markdown
Collaborator

@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, 384

in 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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants