Skip to content

[diffusion] attention: add fp8_fa_sm120 FP8 backend for SM120 GPUs - #40175

Merged
BBuf merged 4 commits into
sgl-project:mainfrom
emre570:fp8-fa-sm120-attention
Sep 22, 2026
Merged

BBuf merged 4 commits into
sgl-project:mainfrom
emre570:fp8-fa-sm120-attention

Conversation

@emre570

@emre570 emre570 commented Sep 18, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

SM120 GPUs (GeForce RTX 50, RTX PRO 6000 Blackwell) have FP8 tensor cores, but the diffusion
runtime has no FP8 attention path for them. On SM120 the default is Torch SDPA, and BF16
attention there is already near the hardware ceiling: cuDNN SDPA reaches 95% of the measured
BF16 mma.sync peak on an RTX 5080. The only way to go faster on these cards is lower precision.

On MiniMax-H3 (56 heads, head_dim 128, S=30272) attention is about 47% of the GPU time of a
denoise step, so this is where an FP8 kernel pays off.

This PR adds fp8_fa_sm120: an opt-in attention backend around an FP8 (E4M3) flash-attention
forward written in CuTe-DSL for SM120.

Modifications

  • python/sglang/kernels/ops/attention/fp8_fa_sm120/
    • kernel.py: the CuTe-DSL kernel. E4M3 QK and PV with FP32 accumulation, online softmax in
      registers, two-stage K/V prefetch so the copy overlaps the math, probabilities packed to
      E4M3 inside registers (no shared-memory round trip), 128 x 32 CTA tile.
    • fused_prep.py: two Triton passes that quantize strided BF16 Q/K/V views per head straight
      into the kernel buffers (amax, then pack). No FP32 intermediates, no layout copy.
    • plan.py: one compiled kernel plus its buffers per (S, H, strides, scale, device).
      bind_inputs() rebinds storage, so all DiT layers share one compilation.
  • multimodal_gen/runtime/layers/attention/backends/fp8_fa_sm120_attn.py: the backend. Handles
    forward and forward_varlen (including the MiniMax-H3 trailing-padding packing). Falls back
    to cuDNN SDPA for causal, batch > 1, head_dim != 128, non-BF16 and non-SM120 calls, and logs
    each fallback reason once.
  • AttentionBackendEnum.FP8_FA_SM120 and a resolver in platforms/cuda.py. Nothing selects the
    backend automatically.
  • Docs: two rows in sglang-diffusion/attention_backends.mdx.
  • Tests: multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py (skipped unless an SM120 GPU is
    present).

No new dependency: nvidia-cutlass-dsl and Triton are already required.

Limits: dense, non-causal, batch 1, head_dim 128, BF16 inputs. The device gate is SM120 only.
SM121 can get the same kernel later, but I have no SM121 hardware to test on, so it is not part
of this PR. Each new sequence length compiles once (about 10 s). The plan keeps E4M3 copies
of Q/K/V plus the output buffer (about 1.1 GB at S=30272, H=56), shared by all layers.

The kernel structure (tiled mma.sync, online softmax, register accumulators) follows the NVIDIA
CuTe FlashAttention-2 example. The TMA layout, swizzle and K/V scheduling choices were informed by Blake Ledden's CuTe-DSL FlashAttention-2 port for SM120 (NVIDIA/cutlass#3030), which was also the public FP8 reference on this hardware before the FlashInfer kernel landed; thanks to Blake for walking me through its design.

The closest existing FP8 kernel for these GPUs is fmha_v2_prefill_sm120 in FlashInfer. The diffusion runtime does not call it. It is the FP8 reference in the tables below.

Accuracy Tests

Per call, against cuDNN SDPA BF16 on the same inputs (H=56, D=128, normally distributed
synthetic Q/K/V): relative RMS 0.0535 to 0.0539 at S=20440 and S=30272. The FlashInfer FP8 kernel
lands in the same band on the same inputs, so the error is the price of E4M3 Q/K/V, not
something specific to this kernel.

End to end, MiniMax-H3 Ref2VA, 768x768, 107 frames, 50 steps, seed 42, RTX PRO 6000:

clip pair PSNR (dB) SSIM
stock Torch SDPA vs stock cuDNN SDPA (noise floor between two BF16 backends) 22.9 0.82
fp8_fa_sm120 vs stock cuDNN SDPA 19.75 0.707

So the FP8 clip sits about 3 dB below the spread between two stock BF16 backends. Subject,
composition and motion are the same and the clip is sharp. I did not find a visible artifact,
but this is a lossy backend and that is why it is opt-in only. Clips are deterministic across
pods (the cuDNN clip from two different pods is bit-identical).

Unit tests (7) cover: output against cuDNN at two sequence lengths including a masked tail,
plan reuse across calls and across impl instances, H3 trailing-padding varlen, multi-segment
varlen, causal fallback, and resolution by name.

Frames 0, 53 and 106 of the two clips, cuDNN SDPA on top and fp8_fa_sm120 below. The full clips
are attached under the image.

Speed Tests and Profiling

Kernel only, H=56, D=128, non-causal, CUDA events, 20 warmup, 5 x 10 iterations, times in ms.
cuDNN SDPA, Torch SDPA and FA4 have no FP8 path on SM120, so their BF16 kernels are the production
baseline this backend replaces. FlashInfer fmha_v2_prefill_sm120 is the only other FP8 attention
kernel for these GPUs, so it is the FP8 reference.

RTX 5080:

S fp8_fa_sm120 (FP8) FlashInfer fmha_v2_prefill_sm120 (FP8) cuDNN SDPA (BF16) Torch FlashAttention-2 SDPA (BF16) SGLang FA4 SM120 (BF16)
30272 123.6 131.8 226.2 245.4 251.8
20440 57.0 60.4 106.5 112.4 115.5

RTX PRO 6000:

S fp8_fa_sm120 (FP8) FlashInfer fmha_v2_prefill_sm120 (FP8) cuDNN SDPA (BF16) Torch FlashAttention-2 SDPA (BF16)
30272 43.7 53.7 73.7 74.4
20440 19.9 25.0 32.3 34.2

End to end on MiniMax-H3 Ref2VA (768x768, 107 frames, 50 steps, RTX PRO 6000, same pod, v0.5.19
tag with stock pins): 6.87 s/it with torch_cudnn_sdpa, 5.61 s/it with fp8_fa_sm120 (-18.3%).

Checklist


CI States

Latest PR Test (Base): ⏳ Run #35700395424
Latest PR Test (Extra): ⏳ Run #35700395339
Latest PR Test (AMD ROCm 10): ⏳ Run #35700395499

Opt-in attention backend around an FP8 (E4M3) flash-attention forward written in
CuTe-DSL for SM120 (GeForce RTX 50, RTX PRO 6000 Blackwell).

- kernels/ops/attention/fp8_fa_sm120: CuTe-DSL kernel, fused Triton quantization of
  strided BF16 Q/K/V views, and a reusable plan shared by all DiT layers.
- multimodal_gen attention backend with forward and forward_varlen; falls back to
  cuDNN SDPA for causal, batch > 1, head_dim != 128, non-BF16 and non-SM120 calls.
- AttentionBackendEnum.FP8_FA_SM120, CUDA resolver, docs rows and unit tests.

MiniMax-H3 Ref2VA on RTX PRO 6000, 50 steps: 6.87 -> 5.61 s/it against cuDNN SDPA.
query.shape[1],
)
plan = FP8AttentionPlan(query, key, value, self.softmax_scale)
self.plans[key_] = plan

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we bound the GPU memory retained by this cache? Each new sequence length/stride adds a plan permanently to the module-level _PLAN_CACHE. A plan owns the FP8 buffers, output/LSE, and views of the last BF16 Q/K/V through self.inputs, so request completion does not release them.

For S=30272, H=56, D=128, this is roughly 1.02 GiB of plan buffers plus 1.21 GiB of retained Q/K/V per shape (calculated from tensor sizes, not measured). A long-running worker receiving different video lengths/resolutions can therefore run out of memory after several distinct shapes.

Please separate compiled-kernel caching from bounded workspace ownership, release input references when safe, and add a multi-shape memory regression test. Any eviction needs to account for in-flight GPU work and captured graphs.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@BBuf Thanks, I checked it and it's confirmed

They are all fixed:

  • Now the cache holds only the compiled kernel, keyed by (S, H, device), same pattern as cutedsl_gdn.py. Each entry is KBs of host state.
  • The FP8/output/LSE buffers come from the caching allocator on every call and the output belongs to the caller. No input references are kept. Allocation plus the DLPack views cost 47 us on RTX 5080 against 2.9 ms (S=4096) and 127 ms (S=30272) of kernel time, so I did not add a workspace slot.
  • No eviction exists, so there is nothing to synchronize. A freed block is only reused by later work on the same stream, and allocation, prep and launch all run on the current stream inside one call. Attention runs eagerly under the breakable CUDA graph, so no captured graph holds these addresses.
  • Per-call torch.empty exposed a second problem. The prep left unwritten positions inside the last 16-key group of V^T (the key permutation), which was only safe because the plan zeroed its buffers once. The pack pass now writes every padding position.
  • The tests test_distinct_shapes_release_memory (memory returns to baseline across sequence lengths) and test_prep_defines_every_padded_byte (no NaN byte survives the prep in a 0xFF-filled workspace) both fail on the previous commit. Output and LSE are bit-identical to the previous commit for S in (1052, 4096, 5980, 30272).

…ernels only

_PLAN_CACHE kept one plan per (S, H, strides, scale) forever. Each plan owned about
1.0 GiB of FP8, output and LSE buffers at S=30272 H=56 and a reference to the last
BF16 Q/K/V views (1.2 GiB), so a worker serving several sequence lengths ran out of
memory.

- plan.py: fp8_attention() replaces FP8AttentionPlan. The compiled kernel is cached
  per (S, H, device); the workspace comes from the caching allocator on every call
  and the output belongs to the caller. No input references are kept. Allocation
  plus DLPack views cost 47 us on an RTX 5080 against 2.9 ms (S=4096) and 127 ms
  (S=30272) of kernel time.
- fused_prep.py: the pack pass writes every position of the padded buffers (stores
  masked on the buffer position, grid up to padded_queries), so the workspace can
  come from torch.empty. The V^T key permutation left unwritten positions inside
  the last 16-key group; they were only safe because the plan zeroed its buffers
  once.
- Tests: memory returns to baseline across distinct sequence lengths; no NaN byte
  survives the prep in a 0xFF-filled workspace.

Output and LSE are bit-identical to the previous commit for S in (1052, 4096, 5980,
30272), H in (8, 56).
@BBuf
BBuf requested a review from kevin-mii as a code owner September 22, 2026 00:18
@BBuf BBuf added run-ci CI: run the baseline test suite on this PR run-ci-extra CI: also run the extra suite (requires run-ci) labels Sep 22, 2026
@BBuf

BBuf commented Sep 22, 2026

Copy link
Copy Markdown
Collaborator

@emre570 Please fix lint

@emre570

emre570 commented Sep 22, 2026

Copy link
Copy Markdown
Contributor Author

@emre570 Please fix lint

@BBuf everything fixed, there should be no problem

@BBuf
BBuf merged commit 9d58189 into sgl-project:main Sep 22, 2026
145 of 182 checks passed
detain added a commit to detain/sglang that referenced this pull request Sep 28, 2026
Brings in 113 upstream commits (7ad55e4..d34f7b2). Upstream now
contains three stack PRs as merged: sgl-project#40501, sgl-project#40175, sgl-project#33778; every line
their squashes add is already present in the stack.

Conflicts resolved:
- arg_groups/fields/memory.py, managers/cache_controller.py: union of the
  stack's fast_file backend (sgl-project#39880) and upstream's tensorcast backend.
- layers/attention/qsa/mqa.py: keep the stack's SM120 Triton imports and
  fp8 _scoring_dtype path; adopt upstream's ROCm 16-wide head alignment
  (sgl-project#38875) and its is_hip import.
- models/qwen4_exp.py: all six hunks are stack-only additions (PLE host
  staging, CUDA-graph prewarm, replay prepare) against upstream's final
  sgl-project#40501; kept ours. The merged file equals the pre-merge stack version.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion SGLang Diffusion documentation Improvements or additions to documentation jit-kernel run-ci CI: run the baseline test suite on this PR run-ci-extra CI: also run the extra suite (requires run-ci)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants