[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
|
# 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
|
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
|
Nothing is trained and no weights change: a cheap estimator runs before
|
||||||
attention and hands the selected kernel a `q2k_block_index`.
|
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}'
|
"min_seq_len": 4096}'
|
||||||
```
|
```
|
||||||
|
|
||||||
**`text_encoder=fa` is not optional.** `--attention-backend` applies to every
|
**The text-encoder override is not optional.** `--attention-backend` applies to
|
||||||
component, and the Qwen3-VL text encoder admits only `fa` / `torch_sdpa` /
|
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
|
or `sage_attn_3`; without the override it raises and the server never starts.
|
||||||
the override on the *encoder*, not the DiT — `transformer=subblock_sparse_attn`
|
Put the override on the *encoder*, not the DiT —
|
||||||
appears to work and silently does nothing, because H3 resolves the DiT backend
|
`transformer=subblock_sparse_attn` appears to work and silently does nothing,
|
||||||
lazily on the first forward, outside the component-loading context that the
|
because H3 resolves the DiT backend lazily on the first forward, outside the
|
||||||
override applies to.
|
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,
|
`--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
|
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 |
|
| dtype | bfloat16 |
|
||||||
| head_dim | 128 |
|
| head_dim | 128 |
|
||||||
| attention | non-causal, one contiguous sequence per call |
|
| 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.
|
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
|
**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
|
resolver accepts exactly compute capability 9.0, 10.0, or 12.0 before loading
|
||||||
kernel, so a B300 or an SM12x GPU fails at launch rather than after ten dense
|
the selected kernel, so a B300 or another unsupported capability fails at
|
||||||
denoise steps. The exact 10.0 check is required because FlashInfer's kernel is
|
launch rather than after ten dense denoise steps. The exact 10.0 check is
|
||||||
built for `sm_100a` and has no forward-compatible 10.3 cubin.
|
required because FlashInfer's kernel is built for `sm_100a` and has no
|
||||||
|
forward-compatible 10.3 cubin.
|
||||||
|
|
||||||
## How the score works
|
## 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
|
``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
|
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
|
CuTe-DSL block-sparse FlashAttention or FlashInfer's architecture-specific
|
||||||
``bsa_attn_blk64_fwd`` (bf16, head_dim 128). The estimator and the measurements
|
SM100/SM120 blk64 kernels (bf16, head_dim 128). The estimator and the
|
||||||
behind its defaults are documented there.
|
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
|
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
|
LOG2E = 1.4426950408889634
|
||||||
BLOCK = 64 # the kernel's block granularity (kSparseBlockSize=64)
|
BLOCK = 64 # the kernel's block granularity (kSparseBlockSize=64)
|
||||||
BUDGET_GRANULARITY = 8 # blocks per query row the kernel bills in, padding to fit
|
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.
|
"""SubBlock block-sparse attention backend.
|
||||||
|
|
||||||
Routes the same 64-token SubBlock plan to SGLang's CuTe-DSL block-sparse
|
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
|
FlashAttention kernel on SM90 or FlashInfer's architecture-specific kernels on
|
||||||
over query/key sub-block pairs selects the blocks (see ``backends/subblock_sparse/``).
|
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
|
Everything is training-free: the router runs before attention and produces
|
||||||
the ``q2k_block_index`` the selected kernel consumes.
|
the ``q2k_block_index`` the selected kernel consumes.
|
||||||
|
|
||||||
@@ -18,15 +19,16 @@ individual keys of the defaults below::
|
|||||||
--attention-backend-config '{"sparsity": 0.85}'
|
--attention-backend-config '{"sparsity": 0.85}'
|
||||||
|
|
||||||
Requirements inherited from the kernels: compute capability 9.0 (Hopper) or
|
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
|
10.0/12.0 (Blackwell), bf16, head_dim 128. Hopper uses SGLang's CuTe-DSL SM90
|
||||||
block-sparse FlashAttention kernel; B200 uses FlashInfer's ``sm_100a`` blk64
|
block-sparse FlashAttention kernel; B200 and SM120 devices use FlashInfer's
|
||||||
kernel. Inside the DiT, any call the kernels cannot serve -- cross/refiner
|
architecture-specific blk64 kernels. Inside the DiT, any call the kernels cannot
|
||||||
attention, short sequences, non-bf16 -- runs dense instead. On any other GPU
|
serve -- cross/refiner attention, short sequences, non-bf16 -- runs dense instead.
|
||||||
the resolver refuses the backend at startup rather than falling back.
|
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
|
``--attention-backend`` reaches every component, and the text encoder admits
|
||||||
only fa / torch_sdpa / sage_attn_3, so pair it with
|
only fa / torch_sdpa / sage_attn_3. Pair it with
|
||||||
``--component-attention-backends text_encoder=fa``; see the README.
|
``--component-attention-backends text_encoder=fa`` on SM90/SM100, or
|
||||||
|
``text_encoder=torch_sdpa`` on SM120; see the README.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
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 (
|
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import (
|
||||||
SubBlockRouter,
|
SubBlockRouter,
|
||||||
load_bsa_attn_blk64_fwd,
|
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.managers.forward_context import get_forward_context
|
||||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
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
|
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)
|
@functools.lru_cache(maxsize=None)
|
||||||
def _get_subblock_sparse_attention_runner(device: torch.device):
|
def _get_subblock_sparse_attention_runner(device: torch.device):
|
||||||
"""Resolve the architecture-specific kernel once per CUDA 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
|
return _sm90_sparse_attention
|
||||||
if capability == (10, 0):
|
if capability == (10, 0):
|
||||||
return _sm100_sparse_attention
|
return _sm100_sparse_attention
|
||||||
|
if capability == (12, 0):
|
||||||
|
return _sm120_sparse_attention
|
||||||
raise RuntimeError(
|
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."
|
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.
|
"""Dispatch a prepared 64x64 routing plan to Hopper or Blackwell.
|
||||||
|
|
||||||
SM90 requires every active index prefix to be sorted in ascending order;
|
SM90 requires every active index prefix to be sorted in ascending order;
|
||||||
SM100 accepts the router's original order. Heterogeneous callers must sort
|
SM100 and SM120 accept the router's original order. Heterogeneous callers
|
||||||
compact sparse prefixes before expanding them to full-width dense rows.
|
must sort compact sparse prefixes before expanding them to full-width dense
|
||||||
|
rows.
|
||||||
"""
|
"""
|
||||||
runner = _get_subblock_sparse_attention_runner(q.device)
|
runner = _get_subblock_sparse_attention_runner(q.device)
|
||||||
return runner(
|
return runner(
|
||||||
|
|||||||
@@ -381,10 +381,10 @@ class _VMOBAAttentionBackendResolver(_CudaAttentionBackendResolver):
|
|||||||
class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
|
class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
|
||||||
backend = AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN
|
backend = AttentionBackendEnum.SUBBLOCK_SPARSE_ATTN
|
||||||
|
|
||||||
# Hopper uses SGLang's SM90 CuTe-DSL block-sparse kernel. Blackwell uses the
|
# Hopper uses SGLang's SM90 CuTe-DSL block-sparse kernel. SM100 uses
|
||||||
# FlashInfer blk64 kernel built specifically for sm_100a; 10.3 and 12.x do
|
# FlashInfer's architecture-specific sm_100a kernel; SM120 uses FlashInfer's
|
||||||
# not have a compatible cubin and must still fail closed.
|
# CuTe-DSL SM120 blk64 kernel. Other capabilities still fail closed.
|
||||||
supported_capabilities = {(9, 0), (10, 0)}
|
supported_capabilities = {(9, 0), (10, 0), (12, 0)}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def resolve(cls, platform) -> str:
|
def resolve(cls, platform) -> str:
|
||||||
@@ -395,8 +395,8 @@ class _SubBlockSparseAttentionBackendResolver(_CudaAttentionBackendResolver):
|
|||||||
if capability_tuple not in cls.supported_capabilities:
|
if capability_tuple not in cls.supported_capabilities:
|
||||||
found = capability.as_version_str() if capability else "unknown"
|
found = capability.as_version_str() if capability else "unknown"
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"SubBlock sparse attention needs compute capability 9.0 "
|
"SubBlock sparse attention needs compute capability 9.0, 10.0, "
|
||||||
f"(Hopper) or 10.0 (B200 / GB200); this device reports {found}."
|
f"or 12.0; this device reports {found}."
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn import ( # noqa: F401
|
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
|
from sglang.kernels.ops.attention.flash_attn.cute.interface import ( # noqa: F401
|
||||||
flash_attn_func,
|
flash_attn_func,
|
||||||
)
|
)
|
||||||
else:
|
elif capability_tuple == (10, 0):
|
||||||
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( # noqa: F401
|
from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse import ( # noqa: F401
|
||||||
load_bsa_attn_blk64_fwd,
|
load_bsa_attn_blk64_fwd,
|
||||||
)
|
)
|
||||||
|
|
||||||
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"
|
return "sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn.SubBlockSparseAttentionBackend"
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Failed to import SubBlock sparse attention: %s", str(e))
|
logger.error("Failed to import SubBlock sparse attention: %s", str(e))
|
||||||
dependency = (
|
dependency = (
|
||||||
"SGLang's SM90 CuTe-DSL FlashAttention dependencies"
|
"SGLang's SM90 CuTe-DSL FlashAttention dependencies"
|
||||||
if capability_tuple == (9, 0)
|
if capability_tuple == (9, 0)
|
||||||
else "FlashInfer with the blk64 block-sparse kernel "
|
else (
|
||||||
|
"FlashInfer with the SM100 blk64 block-sparse kernel "
|
||||||
"(flashinfer.cute_dsl.sparse.bsa_attn_blk64_fwd)"
|
"(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
|
raise ImportError(f"SubBlock sparse attention needs {dependency}.") from e
|
||||||
|
|
||||||
|
|||||||
@@ -2,8 +2,9 @@
|
|||||||
"""SubBlock block-sparse attention backend.
|
"""SubBlock block-sparse attention backend.
|
||||||
|
|
||||||
The schedule and adapter tests are pure CPU. The numerical tests need either
|
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
|
an SM90 GPU with SGLang's CuTe-DSL dependencies, an SM100 GPU with
|
||||||
FlashInfer's ``bsa_attn_blk64_fwd`` and are skipped otherwise.
|
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
|
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
|
``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()
|
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:
|
else:
|
||||||
return False
|
return False
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -60,7 +67,8 @@ def _subblock_kernel_available() -> bool:
|
|||||||
|
|
||||||
|
|
||||||
requires_subblock_kernel = unittest.skipUnless(
|
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.skipTest("requires the SM100 SubBlock kernel")
|
||||||
self._assert_kernel_backed_mixed_query_mask()
|
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):
|
def test_skipped_step_is_bitwise_dense(self):
|
||||||
device = torch.device("cuda")
|
device = torch.device("cuda")
|
||||||
q, k, v = _structured_qkv(self.seq_len, device)
|
q, k, v = _structured_qkv(self.seq_len, device)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_att
|
|||||||
_get_subblock_sparse_attention_runner,
|
_get_subblock_sparse_attention_runner,
|
||||||
_sm90_sparse_attention,
|
_sm90_sparse_attention,
|
||||||
_sm100_sparse_attention,
|
_sm100_sparse_attention,
|
||||||
|
_sm120_sparse_attention,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import (
|
from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import (
|
||||||
_minimax_h3_attention_core_impl,
|
_minimax_h3_attention_core_impl,
|
||||||
@@ -36,6 +37,9 @@ from sglang.multimodal_gen.runtime.platforms import (
|
|||||||
AttentionBackendEnum,
|
AttentionBackendEnum,
|
||||||
current_platform,
|
current_platform,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms.cuda import (
|
||||||
|
_SubBlockSparseAttentionBackendResolver,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -66,12 +70,76 @@ class TestSubBlockSparseAttentionDispatch(CustomTestCase):
|
|||||||
|
|
||||||
self.assertIs(runner, _sm100_sparse_attention)
|
self.assertIs(runner, _sm100_sparse_attention)
|
||||||
|
|
||||||
|
def test_dispatches_sm120(self):
|
||||||
|
device = torch.device("cuda:0")
|
||||||
|
with patch("torch.cuda.get_device_capability", return_value=(12, 0)):
|
||||||
|
runner = _get_subblock_sparse_attention_runner(device)
|
||||||
|
|
||||||
|
self.assertIs(runner, _sm120_sparse_attention)
|
||||||
|
|
||||||
|
def test_platform_resolver_loads_sm120_dependency(self):
|
||||||
|
capability = Mock(major=12, minor=0)
|
||||||
|
capability.as_version_str.return_value = "12.0"
|
||||||
|
platform = Mock()
|
||||||
|
platform.get_device_capability.return_value = capability
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.backends."
|
||||||
|
"subblock_sparse.load_bsa_attn_sm120_blk64_fwd"
|
||||||
|
) as load_sm120:
|
||||||
|
resolved = _SubBlockSparseAttentionBackendResolver.resolve(platform)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
resolved,
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.backends."
|
||||||
|
"subblock_sparse_attn.SubBlockSparseAttentionBackend",
|
||||||
|
)
|
||||||
|
load_sm120.assert_called_once_with()
|
||||||
|
|
||||||
|
def test_sm120_adapter_forwards_subblock_plan(self):
|
||||||
|
q = torch.empty((1, 64, 2, 128), dtype=torch.bfloat16)
|
||||||
|
k = torch.empty((1, 65, 2, 128), dtype=torch.bfloat16)
|
||||||
|
v = torch.empty_like(k)
|
||||||
|
q2k_block_index = torch.zeros((1, 2, 1, 2), dtype=torch.int32)
|
||||||
|
block_counts = torch.tensor([[[2], [1]]], dtype=torch.int32)
|
||||||
|
expected = torch.empty_like(q)
|
||||||
|
kernel = Mock(return_value=(expected, None))
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.backends."
|
||||||
|
"subblock_sparse_attn.load_bsa_attn_sm120_blk64_fwd",
|
||||||
|
return_value=kernel,
|
||||||
|
):
|
||||||
|
result = _sm120_sparse_attention(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
q2k_block_index,
|
||||||
|
topk=2,
|
||||||
|
softmax_scale=0.125,
|
||||||
|
block_counts=block_counts,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(result, expected)
|
||||||
|
kernel.assert_called_once()
|
||||||
|
args, kwargs = kernel.call_args
|
||||||
|
self.assertIs(args[0], q)
|
||||||
|
self.assertIs(args[1], k)
|
||||||
|
self.assertIs(args[2], v)
|
||||||
|
self.assertIs(args[3], q2k_block_index)
|
||||||
|
self.assertEqual(args[4], 2)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
kwargs["block_sizes"], torch.tensor([64, 1], dtype=torch.int32)
|
||||||
|
)
|
||||||
|
self.assertIs(kwargs["q2k_block_nums"], block_counts)
|
||||||
|
self.assertEqual(kwargs["softmax_scale"], 0.125)
|
||||||
|
|
||||||
def test_rejects_unsupported_compute_capability(self):
|
def test_rejects_unsupported_compute_capability(self):
|
||||||
device = torch.device("cuda:0")
|
device = torch.device("cuda:0")
|
||||||
with patch("torch.cuda.get_device_capability", return_value=(10, 3)):
|
with patch("torch.cuda.get_device_capability", return_value=(10, 3)):
|
||||||
with self.assertRaisesRegex(
|
with self.assertRaisesRegex(
|
||||||
RuntimeError,
|
RuntimeError,
|
||||||
"supports compute capability 9.0 or 10.0;.*10.3 device",
|
"supports compute capability 9.0, 10.0, or 12.0;.*10.3 device",
|
||||||
):
|
):
|
||||||
_get_subblock_sparse_attention_runner(device)
|
_get_subblock_sparse_attention_runner(device)
|
||||||
|
|
||||||
@@ -352,6 +420,7 @@ class TestSubBlockSparseAttentionModalities(CustomTestCase):
|
|||||||
for runner, sparse_rows in (
|
for runner, sparse_rows in (
|
||||||
(_sm90_sparse_attention, ([1, 4, 7], [0, 3, 5])),
|
(_sm90_sparse_attention, ([1, 4, 7], [0, 3, 5])),
|
||||||
(_sm100_sparse_attention, ([7, 1, 4], [5, 0, 3])),
|
(_sm100_sparse_attention, ([7, 1, 4], [5, 0, 3])),
|
||||||
|
(_sm120_sparse_attention, ([7, 1, 4], [5, 0, 3])),
|
||||||
):
|
):
|
||||||
with (
|
with (
|
||||||
self.subTest(runner=runner.__name__),
|
self.subTest(runner=runner.__name__),
|
||||||
|
|||||||
Reference in New Issue
Block a user