From 51470b376fbd6059e5998359cfe02683a203a11f Mon Sep 17 00:00:00 2001
From: WenhaoZhang <42087078+niehen6174@users.noreply.github.com>
Date: Sun, 9 Aug 2026 16:26:06 +0800
Subject: [PATCH] [diffusion] feat: support sol-attn sparse attention backend
for h3 (#33702)
---
.../sglang-diffusion/attention_backends.mdx | 78 ++++++
.../layers/attention/backends/sol_attn.py | 240 ++++++++++++++++++
.../runtime/models/dits/minimax_h3.py | 2 +
.../multimodal_gen/runtime/platforms/cuda.py | 25 ++
.../runtime/platforms/interface.py | 2 +
.../test/unit/test_sol_attn_backend.py | 118 +++++++++
6 files changed, 465 insertions(+)
create mode 100644 python/sglang/multimodal_gen/runtime/layers/attention/backends/sol_attn.py
create mode 100644 python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py
diff --git a/docs/docs/sglang-diffusion/attention_backends.mdx b/docs/docs/sglang-diffusion/attention_backends.mdx
index e402a5b90..0e750f8d9 100644
--- a/docs/docs/sglang-diffusion/attention_backends.mdx
+++ b/docs/docs/sglang-diffusion/attention_backends.mdx
@@ -64,6 +64,11 @@ For SGLang-native pipelines, the CLI accepts the lowercase names of `AttentionBa
`SAGE_ATTN_3` |
Requires SageAttention3 installed per upstream instructions. |
+
+ | `sol_attn` |
+ `SOL_ATTN` |
+ Requires the upstream sol-attn package. Install with pip install git+https://github.com/NVlabs/Sana.git@sol-engine#subdirectory=techniques/sparse_backends. BF16, head dim 128. Configure via --attention-backend-config. |
+
| `video_sparse_attn` |
`VIDEO_SPARSE_ATTN` |
@@ -339,6 +344,69 @@ Some backends require additional configuration. You can pass these parameters vi
+**Sol-Attn (`sol_attn`)**
+
+
+
+
+
+
+
+
+
+
+ | Parameter |
+ Type |
+ Description |
+ Default |
+
+
+
+
+ | `tau` |
+ `float` |
+ Routing threshold scale. Higher values select fewer exact KV blocks. |
+ `1.0` |
+
+
+ | `thresh_type` |
+ `str` |
+ Threshold mode: `diag` or `exact`. |
+ `diag` |
+
+
+ | `sink_tokens` |
+ `int` |
+ Exact KV sink length for prefix tokens such as text/audio rows. |
+ `0` |
+
+
+ | `sink_start` |
+ `int` |
+ Start index of the exact KV sink range. |
+ `0` |
+
+
+ | `dense_steps` |
+ `int` |
+ Use dense attention for the first N denoising steps. |
+ `10` |
+
+
+ | `dense_layers` |
+ `str` |
+ Layer indices kept dense, e.g. `0,1` or `0-2`. |
+ `0,1` |
+
+
+ | `kv_splits` |
+ `int | str` |
+ KV split factor passed to the Sol-Attn kernel. Use `auto` on long sequences. |
+ `auto` |
+
+
+
+
## Platform support matrix
@@ -413,6 +481,16 @@ Some backends require additional configuration. You can pass these parameters vi
| ❌ |
CUDA-only (optional dependency). |
+
+ | `sol_attn` |
+ Yes |
+ No |
+ No |
+ No |
+ ❌ |
+ ❌ |
+ CUDA-only. Requires sol-attn. Install with pip install git+https://github.com/NVlabs/Sana.git@sol-engine#subdirectory=techniques/sparse_backends. Configure via --attention-backend-config. |
+
| `video_sparse_attn` |
Yes |
diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sol_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sol_attn.py
new file mode 100644
index 000000000..a63057a87
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sol_attn.py
@@ -0,0 +1,240 @@
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import re
+
+import torch
+
+from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func
+from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
+ AttentionBackend,
+ AttentionImpl,
+ AttentionMetadata,
+)
+from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
+from sglang.multimodal_gen.runtime.server_args import get_global_server_args
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+_SOL_ATTN_HEAD_DIM = 128
+
+
+def _parse_layer_ranges(spec: str | int | None) -> frozenset[int]:
+ if spec is None:
+ return frozenset()
+ if isinstance(spec, int):
+ return frozenset({spec})
+ layers: set[int] = set()
+ for item in str(spec).split(","):
+ item = item.strip()
+ if not item:
+ continue
+ if "-" in item:
+ start, end = item.split("-", 1)
+ layers.update(range(int(start), int(end) + 1))
+ else:
+ layers.add(int(item))
+ return frozenset(layers)
+
+
+def _resolve_kv_splits(q: torch.Tensor, kv_splits: int | str | None) -> int:
+ if kv_splits not in (None, "auto"):
+ return int(kv_splits)
+ arch = tuple(torch.cuda.get_device_capability(q.device))
+ if arch == (9, 0) and q.shape[1] >= 65536:
+ try:
+ import cuda.bindings.driver # noqa: F401
+ import cutlass.cute # noqa: F401
+
+ return 4
+ except ImportError:
+ pass
+ return 1
+
+
+def _get_sol_attn_runtime_config() -> dict:
+ server_args = get_global_server_args()
+ cfg = getattr(server_args, "attention_backend_config", None) or {}
+ dense_layers = cfg.get("dense_layers", "0,1")
+ sink_start = cfg.get("sink_start", 0)
+ return {
+ "tau": float(cfg.get("tau", 1.0)),
+ "thresh_type": str(cfg.get("thresh_type", "diag")),
+ "kv_splits": cfg.get("kv_splits", "auto"),
+ "sink_tokens": int(cfg.get("sink_tokens", 0)),
+ "sink_start": None if sink_start is None else int(sink_start),
+ "dense_steps": int(cfg.get("dense_steps", 10)),
+ "dense_layers": _parse_layer_ranges(dense_layers),
+ }
+
+
+class SolAttnBackend(AttentionBackend):
+ accept_output_buffer: bool = True
+
+ @staticmethod
+ def get_supported_head_sizes() -> list[int]:
+ return [_SOL_ATTN_HEAD_DIM]
+
+ @staticmethod
+ def get_enum() -> AttentionBackendEnum:
+ return AttentionBackendEnum.SOL_ATTN
+
+ @staticmethod
+ def get_impl_cls() -> type[SolAttnImpl]:
+ return SolAttnImpl
+
+
+class SolAttnImpl(AttentionImpl):
+
+ def __init__(
+ self,
+ num_heads: int,
+ head_size: int,
+ causal: bool,
+ softmax_scale: float,
+ num_kv_heads: int | None = None,
+ prefix: str = "",
+ **extra_impl_args,
+ ) -> None:
+ del num_heads, num_kv_heads, extra_impl_args
+ if head_size != _SOL_ATTN_HEAD_DIM:
+ raise ValueError(
+ f"Sol-Attn requires head_size={_SOL_ATTN_HEAD_DIM}, got {head_size}"
+ )
+ self.causal = causal
+ self.softmax_scale = softmax_scale
+ self.prefix = prefix
+ self.layer_idx = self._parse_layer_idx(prefix)
+
+ @staticmethod
+ def _parse_layer_idx(prefix: str) -> int | None:
+ match = re.search(r"blocks\.(\d+)", prefix)
+ if match is None:
+ return None
+ return int(match.group(1))
+
+ def _should_use_dense(self) -> bool:
+ cfg = _get_sol_attn_runtime_config()
+ try:
+ from sglang.multimodal_gen.runtime.managers.forward_context import (
+ get_forward_context,
+ )
+
+ step = int(get_forward_context().current_timestep)
+ except AssertionError:
+ step = 0
+ if step < cfg["dense_steps"]:
+ return True
+ if self.layer_idx is not None and self.layer_idx in cfg["dense_layers"]:
+ return True
+ return False
+
+ def _dense_varlen(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ *,
+ cu_seqlens: torch.Tensor,
+ max_seqlen: int,
+ ) -> torch.Tensor:
+ output = flash_attn_varlen_func(
+ query,
+ key,
+ value,
+ cu_seqlens_q=cu_seqlens,
+ cu_seqlens_k=cu_seqlens,
+ max_seqlen_q=max_seqlen,
+ max_seqlen_k=max_seqlen,
+ softmax_scale=self.softmax_scale,
+ causal=self.causal,
+ )
+ return output[0] if isinstance(output, tuple) else output
+
+ def _run_sol_attn_thd(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ ) -> torch.Tensor:
+ from sol_attn import sol_attn
+
+ cfg = _get_sol_attn_runtime_config()
+ q = query.unsqueeze(0).contiguous()
+ k = key.unsqueeze(0).contiguous()
+ v = value.unsqueeze(0).contiguous()
+ if q.dtype != torch.bfloat16:
+ raise TypeError(f"Sol-Attn requires bfloat16 activations, got {q.dtype}")
+ out = sol_attn(
+ q,
+ k,
+ v,
+ tau=cfg["tau"],
+ thresh_type=cfg["thresh_type"],
+ kv_splits=_resolve_kv_splits(q, cfg["kv_splits"]),
+ sink_start=cfg["sink_start"],
+ sink_tokens=cfg["sink_tokens"],
+ )
+ return out.squeeze(0)
+
+ def forward(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ attn_metadata: AttentionMetadata,
+ ) -> torch.Tensor:
+ del attn_metadata
+ if self._should_use_dense():
+ q = query.transpose(1, 2).reshape(
+ query.shape[0] * query.shape[1], query.shape[2], query.shape[3]
+ )
+ k = key.transpose(1, 2).reshape(
+ key.shape[0] * key.shape[1], key.shape[2], key.shape[3]
+ )
+ v = value.transpose(1, 2).reshape(
+ value.shape[0] * value.shape[1], value.shape[2], value.shape[3]
+ )
+ cu_seqlens = torch.arange(
+ 0,
+ (query.shape[0] + 1) * query.shape[1],
+ query.shape[1],
+ device=query.device,
+ dtype=torch.int32,
+ )
+ out = self._dense_varlen(
+ q,
+ k,
+ v,
+ cu_seqlens=cu_seqlens,
+ max_seqlen=query.shape[1],
+ )
+ return out.reshape(query.shape[0], query.shape[1], query.shape[2], -1)
+ q = query.reshape(query.shape[0] * query.shape[1], query.shape[2], -1)
+ k = key.reshape(key.shape[0] * key.shape[1], key.shape[2], -1)
+ v = value.reshape(value.shape[0] * value.shape[1], value.shape[2], -1)
+ out = self._run_sol_attn_thd(q, k, v)
+ return out.reshape(query.shape[0], query.shape[1], query.shape[2], -1)
+
+ def forward_varlen(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ *,
+ cu_seqlens: torch.Tensor,
+ max_seqlen: int,
+ cu_seqlens_host: tuple[int, ...] | None = None,
+ ) -> torch.Tensor:
+ del cu_seqlens_host
+ if self._should_use_dense():
+ return self._dense_varlen(
+ query,
+ key,
+ value,
+ cu_seqlens=cu_seqlens,
+ max_seqlen=max_seqlen,
+ )
+ return self._run_sol_attn_thd(query, key, value)
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py
index 25b5ae31b..d9524a0c2 100644
--- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py
@@ -527,6 +527,7 @@ class MiniMaxH3Attention(nn.Module):
self.inner_dim = self.total_num_heads * self.head_dim
self.local_inner_dim = self.num_heads * self.head_dim
self.softmax_scale = self.head_dim**-0.5
+ self.prefix = prefix
self._attention_impl = None
self._attention_backend_enum: AttentionBackendEnum | None = None
# The checkpoint stores one fused qkv tensor. Each logical Q/K/V
@@ -576,6 +577,7 @@ class MiniMaxH3Attention(nn.Module):
causal=False,
softmax_scale=self.softmax_scale,
num_kv_heads=self.num_heads,
+ prefix=self.prefix,
)
# Ring only supports FA (see _minimax_h3_attention_core_impl); keep
# the resolved enum alongside the impl instance instead of a second
diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py
index 9000c22e6..e2b4518fa 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py
@@ -255,6 +255,30 @@ class _SparseVideoGen2AttentionBackendResolver(_CudaAttentionBackendResolver):
) from e
+class _SolAttnBackendResolver(_CudaAttentionBackendResolver):
+ backend = AttentionBackendEnum.SOL_ATTN
+
+ @classmethod
+ def resolve(cls, platform) -> str:
+ try:
+ from sol_attn import sol_attn # noqa: F401
+
+ from sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn import ( # noqa: F401
+ SolAttnBackend,
+ )
+
+ return (
+ "sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn."
+ "SolAttnBackend"
+ )
+ except ImportError as e:
+ logger.error("Failed to import Sol-Attn backend: %s", str(e))
+ raise ImportError(
+ "Sol-Attn backend is not installed. Install it with "
+ "`pip install git+https://github.com/NVlabs/Sana.git@sol-engine#subdirectory=techniques/sparse_backends`."
+ ) from e
+
+
class _VMOBAAttentionBackendResolver(_CudaAttentionBackendResolver):
backend = AttentionBackendEnum.VMOBA_ATTN
@@ -312,6 +336,7 @@ _CUDA_ATTENTION_BACKEND_RESOLVERS = {
_SageAttention3BackendResolver,
_VideoSparseAttentionBackendResolver,
_SparseVideoGen2AttentionBackendResolver,
+ _SolAttnBackendResolver,
_VMOBAAttentionBackendResolver,
_FlashAttention2BackendResolver,
_FlashAttentionBackendResolver,
diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py
index fc51fa743..2bd5d7a1f 100644
--- a/python/sglang/multimodal_gen/runtime/platforms/interface.py
+++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py
@@ -43,6 +43,7 @@ class AttentionBackendEnum(enum.Enum):
LASER_ATTN = enum.auto()
BLOCK_SPARSE_ATTN = enum.auto()
RAIN_FUSION_ATTN = enum.auto()
+ SOL_ATTN = enum.auto()
NO_ATTENTION = enum.auto()
def __str__(self):
@@ -60,6 +61,7 @@ class AttentionBackendEnum(enum.Enum):
AttentionBackendEnum.LASER_ATTN,
AttentionBackendEnum.BLOCK_SPARSE_ATTN,
AttentionBackendEnum.RAIN_FUSION_ATTN,
+ AttentionBackendEnum.SOL_ATTN,
}
diff --git a/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py b/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py
new file mode 100644
index 000000000..c9466ab33
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py
@@ -0,0 +1,118 @@
+import importlib.util
+import unittest
+from unittest.mock import MagicMock, patch
+
+import torch
+
+from sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn import (
+ SolAttnBackend,
+ SolAttnImpl,
+ _parse_layer_ranges,
+)
+from sglang.multimodal_gen.runtime.platforms.cuda import CudaPlatformBase
+from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum
+
+
+class FakeCudaPlatform(CudaPlatformBase):
+ is_sm120_device = False
+ is_blackwell_device = False
+ supports_flash_attention = True
+
+ @classmethod
+ def is_sm120(cls):
+ return cls.is_sm120_device
+
+ @classmethod
+ def is_blackwell(cls):
+ return cls.is_blackwell_device
+
+ @classmethod
+ def has_device_capability(
+ cls,
+ capability: tuple[int, int] | int,
+ device_id: int = 0,
+ ) -> bool:
+ return cls.supports_flash_attention
+
+
+class TestSolAttnBackend(unittest.TestCase):
+ def test_enum_name(self):
+ self.assertEqual(str(AttentionBackendEnum.SOL_ATTN), "sol_attn")
+ self.assertTrue(AttentionBackendEnum.SOL_ATTN.is_sparse)
+
+ def test_parse_layer_ranges(self):
+ self.assertEqual(_parse_layer_ranges("0,1,3-5"), frozenset({0, 1, 3, 4, 5}))
+
+ def test_backend_head_size(self):
+ self.assertEqual(SolAttnBackend.get_supported_head_sizes(), [128])
+
+ def test_dense_guard_uses_early_steps(self):
+ impl = SolAttnImpl(
+ num_heads=8,
+ head_size=128,
+ causal=False,
+ softmax_scale=128**-0.5,
+ prefix="blocks.5.attn",
+ )
+ ctx = MagicMock()
+ ctx.current_timestep = 3
+ server_args = MagicMock()
+ server_args.attention_backend_config = {
+ "dense_steps": 10,
+ "dense_layers": "0,1",
+ }
+ with (
+ patch(
+ "sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn.get_global_server_args",
+ return_value=server_args,
+ ),
+ patch(
+ "sglang.multimodal_gen.runtime.managers.forward_context.get_forward_context",
+ return_value=ctx,
+ ),
+ ):
+ self.assertTrue(impl._should_use_dense())
+
+ def test_sparse_layer_after_dense_guard(self):
+ impl = SolAttnImpl(
+ num_heads=8,
+ head_size=128,
+ causal=False,
+ softmax_scale=128**-0.5,
+ prefix="blocks.5.attn",
+ )
+ ctx = MagicMock()
+ ctx.current_timestep = 20
+ server_args = MagicMock()
+ server_args.attention_backend_config = {
+ "dense_steps": 10,
+ "dense_layers": "0,1",
+ }
+ with (
+ patch(
+ "sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn.get_global_server_args",
+ return_value=server_args,
+ ),
+ patch(
+ "sglang.multimodal_gen.runtime.managers.forward_context.get_forward_context",
+ return_value=ctx,
+ ),
+ ):
+ self.assertFalse(impl._should_use_dense())
+
+ def test_cuda_resolver(self):
+ if importlib.util.find_spec("sol_attn") is None:
+ self.skipTest("sol_attn package is not available")
+ cls_str = FakeCudaPlatform.get_attn_backend_cls_str(
+ selected_backend=AttentionBackendEnum.SOL_ATTN,
+ head_size=128,
+ dtype=torch.bfloat16,
+ )
+ self.assertTrue(cls_str.endswith("SolAttnBackend"))
+
+ def test_supports_packed_varlen(self):
+ self.assertTrue(SolAttnBackend.supports_packed_varlen())
+
+
+if __name__ == "__main__":
+ unittest.main()