[Diffusion][minimax-h3] Add SM120 support for SubBlock sparse attention (#37332)

Co-authored-by: 全力 <liquanli.lql@antgroup.com>
This commit is contained in:
Quanli Li
2026-09-03 22:05:22 +08:00
committed by GitHub
co-authored by 全力
parent 3239baef25
commit 4e37882a93
7 changed files with 196 additions and 43 deletions
@@ -1,7 +1,8 @@
# SubBlock sparse attention — training-free block sparsity for the MiniMax-H3 DiT
Routes the same 64-token SubBlock plan to SGLang's CuTe-DSL block-sparse
FlashAttention kernel on SM90 or FlashInfer's `bsa_attn_blk64_fwd` on SM100.
FlashAttention kernel on SM90 or FlashInfer's architecture-specific blk64
kernels on SM100 and SM120.
Nothing is trained and no weights change: a cheap estimator runs before
attention and hands the selected kernel a `q2k_block_index`.
@@ -18,13 +19,17 @@ sglang serve --model-path MiniMaxAI/MiniMax-H3 --model-variant fl2va \
"min_seq_len": 4096}'
```
**`text_encoder=fa` is not optional.** `--attention-backend` applies to every
component, and the Qwen3-VL text encoder admits only `fa` / `torch_sdpa` /
`sage_attn_3`; without the override it raises and the server never starts. Put
the override on the *encoder*, not the DiT — `transformer=subblock_sparse_attn`
appears to work and silently does nothing, because H3 resolves the DiT backend
lazily on the first forward, outside the component-loading context that the
override applies to.
**The text-encoder override is not optional.** `--attention-backend` applies to
every component, and the Qwen3-VL text encoder admits only `fa`, `torch_sdpa`,
or `sage_attn_3`; without the override it raises and the server never starts.
Put the override on the *encoder*, not the DiT —
`transformer=subblock_sparse_attn` appears to work and silently does nothing,
because H3 resolves the DiT backend lazily on the first forward, outside the
component-loading context that the override applies to.
On SM120, use `text_encoder=torch_sdpa` instead. The CUDA platform selects
Torch SDPA for dense attention on SM12.x, and component-specific backend
requests are validated strictly.
`--attention-backend-config` is optional and overrides only the keys it names,
so `'{"sparsity": 0.85}'` alone trades quality for another 6%. Inline JSON gets
@@ -38,7 +43,7 @@ are listed below.
| | |
| --- | --- |
| GPU | **compute capability 9.0 or 10.0** — H100 / H200 use SGLang's CuTe-DSL SM90 block-sparse FlashAttention kernel; B200 / GB200 use FlashInfer's architecture-specific `sm_100a` kernel. Other capabilities, including 10.3 (B300 / GB300) and 12.x (RTX PRO 6000, RTX 50xx), are rejected. |
| GPU | **compute capability 9.0, 10.0, or 12.0** — H100 / H200 use SGLang's CuTe-DSL SM90 block-sparse FlashAttention kernel; B200 / GB200 use FlashInfer's architecture-specific `sm_100a` kernel; SM120 devices use FlashInfer's `bsa_attn_sm120_blk64_fwd` CuTe-DSL kernel. Other capabilities, including 10.3 (B300 / GB300), are rejected. |
| dtype | bfloat16 |
| head_dim | 128 |
| attention | non-causal, one contiguous sequence per call |
@@ -48,10 +53,11 @@ refiner, sequences under `min_seq_len`, non-bf16 activations, head_dim != 128
falls back to dense for that call, so no layer has to be excluded by hand.
**On an unsupported GPU it is not a fallback, it is an error at startup.** The
resolver accepts exactly compute capability 9.0 or 10.0 before loading either
kernel, so a B300 or an SM12x GPU fails at launch rather than after ten dense
denoise steps. The exact 10.0 check is required because FlashInfer's kernel is
built for `sm_100a` and has no forward-compatible 10.3 cubin.
resolver accepts exactly compute capability 9.0, 10.0, or 12.0 before loading
the selected kernel, so a B300 or another unsupported capability fails at
launch rather than after ten dense denoise steps. The exact 10.0 check is
required because FlashInfer's kernel is built for `sm_100a` and has no
forward-compatible 10.3 cubin.
## How the score works
@@ -6,11 +6,19 @@ Originally vendored from the standalone SubBlock repository; ``router.py`` and
``router.py`` scores every (query block, key block) pair from sub-block-pooled
Q/K and turns the scores into a ``q2k_block_index`` consumed by SGLang's SM90
CuTe-DSL block-sparse FlashAttention or FlashInfer's SM100
``bsa_attn_blk64_fwd`` (bf16, head_dim 128). The estimator and the measurements
behind its defaults are documented there.
CuTe-DSL block-sparse FlashAttention or FlashInfer's architecture-specific
SM100/SM120 blk64 kernels (bf16, head_dim 128). The estimator and the
measurements behind its defaults are documented there.
"""
from .router import SubBlockRouter, load_bsa_attn_blk64_fwd
from .router import (
SubBlockRouter,
load_bsa_attn_blk64_fwd,
load_bsa_attn_sm120_blk64_fwd,
)
__all__ = ["SubBlockRouter", "load_bsa_attn_blk64_fwd"]
__all__ = [
"SubBlockRouter",
"load_bsa_attn_blk64_fwd",
"load_bsa_attn_sm120_blk64_fwd",
]
@@ -128,6 +128,22 @@ def load_bsa_attn_blk64_fwd():
return mod.bsa_attn_blk64_fwd
@functools.lru_cache(maxsize=1)
def load_bsa_attn_sm120_blk64_fwd():
"""FlashInfer's CuTe-DSL 64-block sparse attention entry point for SM120."""
try:
from flashinfer.cute_dsl.sparse.bsa_attn_sm120 import (
bsa_attn_sm120_blk64_fwd,
)
except Exception as exc:
raise ImportError(
"SM120 SubBlock sparse attention requires FlashInfer's "
"flashinfer.cute_dsl.sparse.bsa_attn_sm120 module"
) from exc
return bsa_attn_sm120_blk64_fwd
LOG2E = 1.4426950408889634
BLOCK = 64 # the kernel's block granularity (kSparseBlockSize=64)
BUDGET_GRANULARITY = 8 # blocks per query row the kernel bills in, padding to fit
@@ -2,8 +2,9 @@
"""SubBlock block-sparse attention backend.
Routes the same 64-token SubBlock plan to SGLang's CuTe-DSL block-sparse
FlashAttention kernel on SM90 or FlashInfer's kernel on SM100. A log-sum-exp
over query/key sub-block pairs selects the blocks (see ``backends/subblock_sparse/``).
FlashAttention kernel on SM90 or FlashInfer's architecture-specific kernels on
SM100 and SM120. A log-sum-exp over query/key sub-block pairs selects the blocks
(see ``backends/subblock_sparse/``).
Everything is training-free: the router runs before attention and produces
the ``q2k_block_index`` the selected kernel consumes.
@@ -18,15 +19,16 @@ individual keys of the defaults below::
--attention-backend-config '{"sparsity": 0.85}'
Requirements inherited from the kernels: compute capability 9.0 (Hopper) or
10.0 (B200 / GB200), bf16, head_dim 128. Hopper uses SGLang's CuTe-DSL SM90
block-sparse FlashAttention kernel; B200 uses FlashInfer's ``sm_100a`` blk64
kernel. Inside the DiT, any call the kernels cannot serve -- cross/refiner
attention, short sequences, non-bf16 -- runs dense instead. On any other GPU
the resolver refuses the backend at startup rather than falling back.
10.0/12.0 (Blackwell), bf16, head_dim 128. Hopper uses SGLang's CuTe-DSL SM90
block-sparse FlashAttention kernel; B200 and SM120 devices use FlashInfer's
architecture-specific blk64 kernels. Inside the DiT, any call the kernels cannot
serve -- cross/refiner attention, short sequences, non-bf16 -- runs dense instead.
On any other GPU the resolver refuses the backend at startup rather than falling back.
``--attention-backend`` reaches every component, and the text encoder admits
only fa / torch_sdpa / sage_attn_3, so pair it with
``--component-attention-backends text_encoder=fa``; see the README.
only fa / torch_sdpa / sage_attn_3. Pair it with
``--component-attention-backends text_encoder=fa`` on SM90/SM100, or
``text_encoder=torch_sdpa`` on SM120; see the README.
"""
from __future__ import annotations
@@ -48,6 +50,7 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import (
SubBlockRouter,
load_bsa_attn_blk64_fwd,
load_bsa_attn_sm120_blk64_fwd,
)
from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
@@ -198,6 +201,30 @@ def _sm100_sparse_attention(
return out[0] if isinstance(out, tuple) else out
def _sm120_sparse_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_block_index: torch.Tensor,
topk: int,
softmax_scale: float,
block_counts: torch.Tensor | None = None,
) -> torch.Tensor:
"""Run a SubBlock routing plan through FlashInfer's SM120 kernel."""
logger.info_once("SubBlock sparse attention kernel active: FlashInfer SM120 blk64")
out = load_bsa_attn_sm120_blk64_fwd()(
q,
k,
v,
q2k_block_index,
topk,
block_sizes=_cached_block_sizes(k.shape[1], k.device),
q2k_block_nums=block_counts,
softmax_scale=softmax_scale,
)
return out[0] if isinstance(out, tuple) else out
@functools.lru_cache(maxsize=None)
def _get_subblock_sparse_attention_runner(device: torch.device):
"""Resolve the architecture-specific kernel once per CUDA device."""
@@ -206,8 +233,10 @@ def _get_subblock_sparse_attention_runner(device: torch.device):
return _sm90_sparse_attention
if capability == (10, 0):
return _sm100_sparse_attention
if capability == (12, 0):
return _sm120_sparse_attention
raise RuntimeError(
"SubBlock sparse attention supports compute capability 9.0 or 10.0; "
"SubBlock sparse attention supports compute capability 9.0, 10.0, or 12.0; "
f"this tensor is on a {capability[0]}.{capability[1]} device."
)
@@ -224,8 +253,9 @@ def _run_subblock_sparse_attention(
"""Dispatch a prepared 64x64 routing plan to Hopper or Blackwell.
SM90 requires every active index prefix to be sorted in ascending order;
SM100 accepts the router's original order. Heterogeneous callers must sort
compact sparse prefixes before expanding them to full-width dense rows.
SM100 and SM120 accept the router's original order. Heterogeneous callers
must sort compact sparse prefixes before expanding them to full-width dense
rows.
"""
runner = _get_subblock_sparse_attention_runner(q.device)
return runner(
@@ -381,10 +381,10 @@ class _VMOBAAttentionBackendResolver(_CudaAttentionBackendResolver):
class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
backend = AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN
# Hopper uses SGLang's SM90 CuTe-DSL block-sparse kernel. Blackwell uses the
# FlashInfer blk64 kernel built specifically for sm_100a; 10.3 and 12.x do
# not have a compatible cubin and must still fail closed.
supported_capabilities = {(9, 0), (10, 0)}
# Hopper uses SGLang's SM90 CuTe-DSL block-sparse kernel. SM100 uses
# FlashInfer's architecture-specific sm_100a kernel; SM120 uses FlashInfer's
# CuTe-DSL SM120 blk64 kernel. Other capabilities still fail closed.
supported_capabilities = {(9, 0), (10, 0), (12, 0)}
@classmethod
def resolve(cls, platform) -> str:
@@ -395,8 +395,8 @@ class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
if capability_tuple not in cls.supported_capabilities:
found = capability.as_version_str() if capability else "unknown"
raise ValueError(
"SubBlock sparse attention needs compute capability 9.0 "
f"(Hopper) or 10.0 (B200 / GB200); this device reports {found}."
"SubBlock sparse attention needs compute capability 9.0, 10.0, "
f"or 12.0; this device reports {found}."
)
try:
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn import ( # noqa: F401
@@ -412,20 +412,31 @@ class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
from sglang.kernels.ops.attention.flash_attn.cute.interface import ( # noqa: F401
flash_attn_func,
)
else:
elif capability_tuple == (10, 0):
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( # noqa: F401
load_bsa_attn_blk64_fwd,
)
load_bsa_attn_blk64_fwd()
else:
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import (
load_bsa_attn_sm120_blk64_fwd,
)
load_bsa_attn_sm120_blk64_fwd()
return "sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn.SubBlockSparseAttentionBackend"
except Exception as e:
logger.error("Failed to import SubBlock sparse attention: %s", str(e))
dependency = (
"SGLang's SM90 CuTe-DSL FlashAttention dependencies"
if capability_tuple == (9, 0)
else "FlashInfer with the blk64 block-sparse kernel "
"(flashinfer.cute_dsl.sparse.bsa_attn_blk64_fwd)"
else (
"FlashInfer with the SM100 blk64 block-sparse kernel "
"(flashinfer.cute_dsl.sparse.bsa_attn_blk64_fwd)"
if capability_tuple == (10, 0)
else "FlashInfer with the SM120 blk64 block-sparse kernel "
"(flashinfer.cute_dsl.sparse.bsa_attn_sm120)"
)
)
raise ImportError(f"SubBlock sparse attention needs {dependency}.") from e
@@ -2,8 +2,9 @@
"""SubBlock block-sparse attention backend.
The schedule and adapter tests are pure CPU. The numerical tests need either
an SM90 GPU with SGLang's CuTe-DSL dependencies or an SM100 GPU with
FlashInfer's ``bsa_attn_blk64_fwd`` and are skipped otherwise.
an SM90 GPU with SGLang's CuTe-DSL dependencies, an SM100 GPU with
FlashInfer's ``bsa_attn_blk64_fwd``, or an SM120 GPU with FlashInfer's
``bsa_attn_sm120_blk64_fwd`` and are skipped otherwise.
The trick that makes the sparse kernel checkable against dense attention: at
``sparsity`` just above 0 every block is inside the budget, so the block-sparse
@@ -52,6 +53,12 @@ def _subblock_kernel_available() -> bool:
)
load_bsa_attn_blk64_fwd()
elif capability == (12, 0):
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import (
load_bsa_attn_sm120_blk64_fwd,
)
load_bsa_attn_sm120_blk64_fwd()
else:
return False
except Exception:
@@ -60,7 +67,8 @@ def _subblock_kernel_available() -> bool:
requires_subblock_kernel = unittest.skipUnless(
_subblock_kernel_available(), "needs an SM90 or SM100 SubBlock attention kernel"
_subblock_kernel_available(),
"needs an SM90, SM100, or SM120 SubBlock attention kernel",
)
@@ -491,6 +499,11 @@ class TestSubBlockNumerics(unittest.TestCase):
self.skipTest("requires the SM100 SubBlock kernel")
self._assert_kernel_backed_mixed_query_mask()
def test_sm120_kernel_backed_mixed_query_mask(self):
if torch.cuda.get_device_capability() != (12, 0):
self.skipTest("requires the SM120 SubBlock kernel")
self._assert_kernel_backed_mixed_query_mask()
def test_skipped_step_is_bitwise_dense(self):
device = torch.device("cuda")
q, k, v = _structured_qkv(self.seq_len, device)