diff --git a/docs_new/docs/advanced_features/attention_backend.mdx b/docs_new/docs/advanced_features/attention_backend.mdx
index 966234b7d..8057f803b 100644
--- a/docs_new/docs/advanced_features/attention_backend.mdx
+++ b/docs_new/docs/advanced_features/attention_backend.mdx
@@ -328,6 +328,8 @@ GDN (Gated Delta Network) is a linear attention mechanism with O(n) complexity,
The GDN linear attention layers have their own kernel backends, selected via `--linear-attn-backend` (default: `triton`). You can override the kernel per phase with `--linear-attn-decode-backend` and `--linear-attn-prefill-backend`.
+On SM100/SM103 with CUDA 13+, SGLang automatically selects FlashInfer for GDN prefill when the per-phase override is unset, the base linear-attention backend is Triton, recurrent state is BF16, key/value head dimensions are 128, dynamic chunking and page-major KV layout are disabled, and `--chunked-prefill-size` is between 1 and 8192. Radix caching may be disabled or use the `no_buffer` strategy; extra-buffer strategies require state checkpoint support.
+
@@ -374,6 +376,12 @@ The GDN linear attention layers have their own kernel backends, selected via `--
| ❌ |
❌ |
+
+ | FlashInfer (CUDA, SM90/SM100/SM103) |
+ ✅ |
+ ✅ |
+ ✅ linear chain; tree falls back to Triton |
+
diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py
index 6bf623a56..76f855fd5 100644
--- a/python/sglang/srt/layers/attention/attention_registry.py
+++ b/python/sglang/srt/layers/attention/attention_registry.py
@@ -274,7 +274,10 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
HybridLinearAttnBackend,
Mamba2AttnBackend,
)
- from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend
+ from sglang.srt.layers.attention.linear.gdn_backend import (
+ GDNAttnBackend,
+ maybe_set_default_flashinfer_gdn_prefill,
+ )
else:
from sglang.srt.hardware_backend.npu.attention.ascend_gdn_backend import (
AscendGDNAttnBackend as GDNAttnBackend,
@@ -287,6 +290,8 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
)
check_environments()
+ if runner.hybrid_gdn_config is not None and not is_npu():
+ maybe_set_default_flashinfer_gdn_prefill(runner)
initialize_linear_attn_config(runner.server_args)
hybrid_backend_cls = HybridLinearAttnBackend
if runner.hybrid_gdn_config is not None:
diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py
index 005893c4e..5e3d9d09d 100644
--- a/python/sglang/srt/layers/attention/linear/gdn_backend.py
+++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py
@@ -55,6 +55,47 @@ elif is_cpu():
fused_gdn_gating = torch.ops.sgl_kernel.fused_gdn_gating_cpu
+def maybe_set_default_flashinfer_gdn_prefill(model_runner: ModelRunner) -> None:
+ """Use FlashInfer for the narrow SM100 GDN prefill domain we validated."""
+ args = model_runner.server_args
+ if (
+ args.linear_attn_prefill_backend is not None
+ or args.linear_attn_backend != "triton"
+ or args.enable_page_major_kv_layout
+ or not is_cuda()
+ or torch.cuda.get_device_capability()[0] != 10
+ ):
+ return
+
+ # Extra-buffer strategies need intermediate state checkpoints.
+ if args.uses_mamba_radix_cache and args.mamba_radix_cache_strategy != "no_buffer":
+ return
+
+ cuda_version = torch.version.cuda
+ chunk_size = args.chunked_prefill_size
+ config = model_runner.hybrid_gdn_config
+ if (
+ cuda_version is None
+ or int(cuda_version.split(".", 1)[0]) < 13
+ or args.enable_dynamic_chunking
+ or chunk_size is None
+ or not 1 <= chunk_size <= 8192
+ or getattr(config, "linear_key_head_dim", None) != 128
+ or getattr(config, "linear_value_head_dim", None) != 128
+ or model_runner.req_to_token_pool.mamba_pool.mamba_cache.temporal.dtype
+ != torch.bfloat16
+ ):
+ return
+
+ from sglang.srt.layers.attention.linear.kernels.gdn_flashinfer import (
+ is_flashinfer_gdn_prefill_available,
+ )
+
+ if is_flashinfer_gdn_prefill_available():
+ args.linear_attn_prefill_backend = "flashinfer"
+ rank0_log("Defaulting SM100 GDN prefill backend to FlashInfer.")
+
+
class GDNKernelDispatcher:
"""Dispatches GDN kernel calls to the appropriate backend per mode."""
@@ -64,6 +105,7 @@ class GDNKernelDispatcher:
prefill_backend: LinearAttnKernelBackend,
):
triton_kernel = TritonGDNKernel()
+ self.tree_verify_kernel = triton_kernel
cutedsl_kernel = None
if decode_backend.is_triton():
@@ -251,7 +293,15 @@ class GDNKernelDispatcher:
query_start_loc: torch.Tensor,
**kwargs,
) -> torch.Tensor:
- return self.verify_kernel.target_verify(
+ # FlashInfer verify supports a linear MTP chain. Tree-shaped drafts
+ # carry parent indices and must use Triton even when decode/prefill use
+ # FlashInfer.
+ verify_kernel = (
+ self.tree_verify_kernel
+ if kwargs.get("retrieve_parent_token") is not None
+ else self.verify_kernel
+ )
+ return verify_kernel.target_verify(
A_log=A_log,
dt_bias=dt_bias,
q=q,
diff --git a/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py b/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py
index 24e716cf9..9ce121f75 100644
--- a/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py
+++ b/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py
@@ -72,6 +72,12 @@ def _get_flashinfer_gdn_kernels():
)
+def is_flashinfer_gdn_prefill_available() -> bool:
+ """Return whether the kernel loader can construct the prefill path."""
+ available, prefill_fn, *_ = _get_flashinfer_gdn_kernels()
+ return bool(available and prefill_fn is not None)
+
+
# ---------------------------------------------------------------------------
# Kernel implementation
# ---------------------------------------------------------------------------
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index c3dbc3cc6..be90a79af 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -2011,7 +2011,7 @@ class ServerArgs:
linear_attn_prefill_backend: A[
Optional[str],
Arg(
- help="Override the kernel backend for linear attention prefill/extend. If not set, uses --linear-attn-backend.",
+ help="Override the kernel backend for linear attention prefill/extend. If not set, uses --linear-attn-backend; compatible SM100 GDN models may automatically select FlashInfer.",
choices=LINEAR_ATTN_KERNEL_BACKEND_CHOICES,
),
] = None
diff --git a/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py
new file mode 100644
index 000000000..687e29a8c
--- /dev/null
+++ b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py
@@ -0,0 +1,174 @@
+import unittest
+from types import SimpleNamespace
+from unittest.mock import MagicMock, patch, sentinel
+
+import torch
+
+from sglang.srt.layers.attention.linear import gdn_backend
+from sglang.srt.layers.attention.linear.gdn_backend import (
+ GDNKernelDispatcher,
+ maybe_set_default_flashinfer_gdn_prefill,
+)
+from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
+from sglang.srt.layers.attention.linear.utils import LinearAttnKernelBackend
+from sglang.test.ci.ci_register import register_cpu_ci
+
+register_cpu_ci(est_time=5, suite="base-a-test-cpu")
+
+
+def make_runner(
+ *,
+ state_dtype=torch.bfloat16,
+ key_dim=128,
+ value_dim=128,
+ **arg_overrides,
+):
+ args = SimpleNamespace(
+ linear_attn_backend="triton",
+ linear_attn_prefill_backend=None,
+ uses_mamba_radix_cache=False,
+ enable_page_major_kv_layout=False,
+ mamba_radix_cache_strategy="no_buffer",
+ enable_dynamic_chunking=False,
+ chunked_prefill_size=8192,
+ )
+ for name, value in arg_overrides.items():
+ setattr(args, name, value)
+
+ return SimpleNamespace(
+ server_args=args,
+ hybrid_gdn_config=SimpleNamespace(
+ linear_key_head_dim=key_dim,
+ linear_value_head_dim=value_dim,
+ ),
+ req_to_token_pool=SimpleNamespace(
+ mamba_pool=SimpleNamespace(
+ mamba_cache=SimpleNamespace(temporal=SimpleNamespace(dtype=state_dtype))
+ )
+ ),
+ )
+
+
+class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
+ def apply_policy(
+ self,
+ runner,
+ *,
+ cuda=True,
+ capability=(10, 0),
+ cuda_version="13.0",
+ flashinfer_available=True,
+ ):
+ with (
+ patch.object(gdn_backend, "is_cuda", return_value=cuda),
+ patch.object(torch.cuda, "get_device_capability", return_value=capability),
+ patch.object(torch.version, "cuda", cuda_version),
+ patch(
+ "sglang.srt.layers.attention.linear.kernels.gdn_flashinfer."
+ "is_flashinfer_gdn_prefill_available",
+ return_value=flashinfer_available,
+ ),
+ ):
+ maybe_set_default_flashinfer_gdn_prefill(runner)
+ return runner.server_args.linear_attn_prefill_backend
+
+ def test_selects_flashinfer_for_supported_sm100_gdn(self):
+ self.assertEqual(self.apply_policy(make_runner()), "flashinfer")
+
+ def test_selects_flashinfer_for_no_buffer_radix_cache(self):
+ runner = make_runner(
+ uses_mamba_radix_cache=True,
+ mamba_radix_cache_strategy="no_buffer",
+ )
+ self.assertEqual(self.apply_policy(runner), "flashinfer")
+
+ def test_preserves_explicit_prefill_override(self):
+ for backend in ("triton", "flashinfer", "cutedsl"):
+ with self.subTest(backend=backend):
+ runner = make_runner(linear_attn_prefill_backend=backend)
+ self.assertEqual(self.apply_policy(runner), backend)
+
+ def test_rejects_unsupported_capability(self):
+ cases = (
+ ("non_cuda", {}, {"cuda": False}),
+ ("hopper", {}, {"capability": (9, 0)}),
+ ("future_sm", {}, {"capability": (12, 0)}),
+ ("cuda_12", {}, {"cuda_version": "12.9"}),
+ ("fp32_state", {"state_dtype": torch.float32}, {}),
+ ("key_dim", {"key_dim": 64}, {}),
+ ("value_dim", {"value_dim": 64}, {}),
+ ("missing_api", {}, {"flashinfer_available": False}),
+ )
+ for name, runner_args, hardware in cases:
+ with self.subTest(name=name):
+ self.assertIsNone(
+ self.apply_policy(make_runner(**runner_args), **hardware)
+ )
+
+ def test_rejects_gdn_config_without_qwen_head_dims(self):
+ runner = make_runner()
+ runner.hybrid_gdn_config = SimpleNamespace()
+ self.assertIsNone(self.apply_policy(runner))
+
+ def test_rejects_unvalidated_runtime_modes(self):
+ cases = (
+ ("non_triton_base", {"linear_attn_backend": "cutedsl"}),
+ ("page_major_kv", {"enable_page_major_kv_layout": True}),
+ (
+ "extra_buffer",
+ {
+ "uses_mamba_radix_cache": True,
+ "mamba_radix_cache_strategy": "extra_buffer",
+ },
+ ),
+ (
+ "extra_buffer_lazy",
+ {
+ "uses_mamba_radix_cache": True,
+ "mamba_radix_cache_strategy": "extra_buffer_lazy",
+ },
+ ),
+ ("dynamic_chunk", {"enable_dynamic_chunking": True}),
+ ("unchunked", {"chunked_prefill_size": -1}),
+ ("unknown_chunk", {"chunked_prefill_size": None}),
+ ("large_chunk", {"chunked_prefill_size": 8193}),
+ )
+ for name, runner_args in cases:
+ with self.subTest(name=name):
+ self.assertIsNone(self.apply_policy(make_runner(**runner_args)))
+
+ def test_tree_verify_uses_triton_kernel(self):
+ flashinfer_kernel = MagicMock(supports_target_verify=True)
+ with (
+ patch.object(gdn_backend, "is_cuda", return_value=True),
+ patch(
+ "sglang.srt.layers.attention.linear.kernels.gdn_flashinfer."
+ "FlashInferGDNKernel",
+ return_value=flashinfer_kernel,
+ ),
+ ):
+ dispatcher = GDNKernelDispatcher(
+ LinearAttnKernelBackend.TRITON,
+ LinearAttnKernelBackend.FLASHINFER,
+ )
+
+ self.assertIsInstance(dispatcher.tree_verify_kernel, TritonGDNKernel)
+
+ tensor = sentinel.tensor
+ with patch.object(
+ dispatcher.tree_verify_kernel, "target_verify"
+ ) as tree_verify:
+ dispatcher.target_verify(
+ *([tensor] * 7),
+ ssm_states=tensor,
+ cache_indices=tensor,
+ query_start_loc=tensor,
+ retrieve_parent_token=sentinel.parent_token,
+ )
+
+ tree_verify.assert_called_once()
+ flashinfer_kernel.target_verify.assert_not_called()
+
+
+if __name__ == "__main__":
+ unittest.main()