[Diffusion][minimax-h3] Add SM120 support for SubBlock sparse attention (#37332)
Co-authored-by: 全力 <liquanli.lql@antgroup.com>
This commit is contained in:
+19
-13
@@ -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
|
||||
|
||||
|
||||
+13
-5
@@ -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",
|
||||
]
|
||||
|
||||
+16
@@ -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
|
||||
|
||||
+42
-12
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user