[diffusion] attention: add fp8_fa_sm120 FP8 backend for SM120 GPUs - #40175
Conversation
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 |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
@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.emptyexposed 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) andtest_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).
|
@emre570 Please fix lint |
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>
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-attentionforward 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 inregisters, 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 straightinto 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. Handlesforwardandforward_varlen(including the MiniMax-H3 trailing-padding packing). Falls backto cuDNN SDPA for causal, batch > 1, head_dim != 128, non-BF16 and non-SM120 calls, and logs
each fallback reason once.
AttentionBackendEnum.FP8_FA_SM120and a resolver inplatforms/cuda.py. Nothing selects thebackend automatically.
sglang-diffusion/attention_backends.mdx.multimodal_gen/test/unit/test_fp8_fa_sm120_attn.py(skipped unless an SM120 GPU ispresent).
No new dependency:
nvidia-cutlass-dsland 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_sm120in 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:
fp8_fa_sm120vs stock cuDNN SDPASo 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_sm120below. The full clipsare 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_sm120is the only other FP8 attentionkernel for these GPUs, so it is the FP8 reference.
RTX 5080:
fp8_fa_sm120(FP8)fmha_v2_prefill_sm120(FP8)RTX PRO 6000:
fp8_fa_sm120(FP8)fmha_v2_prefill_sm120(FP8)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 withfp8_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