[GDN] Support FlashInfer GDN prefill with extra-buffer radix cache (#29735)
This commit is contained in:
@@ -10,11 +10,14 @@ from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.attention_unittest.attention_methods.gdn_attention import (
|
||||
GDNAttentionCase,
|
||||
build_gdn_attention_fixture,
|
||||
make_gdn_cases,
|
||||
run_gdn_attention_case,
|
||||
run_gdn_fixture_eager,
|
||||
)
|
||||
from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner import (
|
||||
run_gdn_cuda_graph_decode_case,
|
||||
@@ -30,6 +33,12 @@ from sglang.test.kits.attention_unittest.runner_modes.split_op_runner import (
|
||||
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
|
||||
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
|
||||
|
||||
_cuda_major = int(torch.version.cuda.split(".")[0]) if torch.version.cuda else 0
|
||||
_sm_major = torch.cuda.get_device_capability()[0] if torch.cuda.is_available() else 0
|
||||
_supports_flashinfer_linear_gdn = _sm_major == 9 or (
|
||||
_sm_major == 10 and _cuda_major >= 13
|
||||
)
|
||||
|
||||
|
||||
@unittest.skipIf(
|
||||
not torch.cuda.is_available() or not is_flashinfer_available(),
|
||||
@@ -322,5 +331,71 @@ class TestFlashInferGDNBackendCorrectness(CustomTestCase):
|
||||
)
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
torch.cuda.is_available()
|
||||
and is_flashinfer_available()
|
||||
and _supports_flashinfer_linear_gdn,
|
||||
"FlashInfer linear GDN requires SM90 or SM100/SM103 with CUDA 13+",
|
||||
)
|
||||
class TestFlashInferLinearGDNBackendCorrectness(CustomTestCase):
|
||||
# FlashInfer's DSL prefill kernels require head size 128 on SM90 and SM100.
|
||||
HEAD_DIM = 128
|
||||
CHECKPOINT_CASE = GDNAttentionCase(
|
||||
name="flashinfer_gdn_prefill_state_checkpoints",
|
||||
backend="triton",
|
||||
linear_attn_prefill_backend="flashinfer",
|
||||
forward_mode=ForwardMode.EXTEND,
|
||||
num_k_heads=2,
|
||||
num_v_heads=4,
|
||||
page_size=16,
|
||||
prefix_lens=(0, 64, 128),
|
||||
extend_lens=(64, 65, 129),
|
||||
)
|
||||
|
||||
def test_prefill_tracked_state_checkpoints(self):
|
||||
fixture = build_gdn_attention_fixture(
|
||||
self,
|
||||
self.CHECKPOINT_CASE,
|
||||
head_k_dim=self.HEAD_DIM,
|
||||
head_v_dim=self.HEAD_DIM,
|
||||
max_context_len=320,
|
||||
runner_batch_size=6,
|
||||
)
|
||||
batch = fixture.forward_batch
|
||||
# Simulate the tracking metadata produced by the extra-buffer scheduler.
|
||||
# This test covers checkpoint mapping and state copies, not scheduler setup.
|
||||
batch.mamba_track_mask = torch.ones(3, dtype=torch.bool, device="cuda")
|
||||
batch.mamba_track_indices = torch.tensor(
|
||||
[4, 5, 6], dtype=torch.int64, device="cuda"
|
||||
)
|
||||
batch.mamba_track_seqlens = torch.tensor(
|
||||
# The final entry selects the second checkpoint at absolute S256.
|
||||
[64, 129, 257],
|
||||
dtype=torch.int64,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
cache = fixture.runner.req_to_token_pool.mamba2_layer_cache(0)
|
||||
initial_conv = cache.conv[0].clone()
|
||||
initial_ssm = cache.temporal.clone()
|
||||
flashinfer_output = run_gdn_fixture_eager(fixture)
|
||||
flashinfer_tracked = cache.temporal[batch.mamba_track_indices].clone()
|
||||
|
||||
cache.conv[0].copy_(initial_conv)
|
||||
cache.temporal.copy_(initial_ssm)
|
||||
fixture.backend.linear_attn_backend.kernel_dispatcher.extend_kernel = (
|
||||
TritonGDNKernel()
|
||||
)
|
||||
triton_output = run_gdn_fixture_eager(fixture)
|
||||
triton_tracked = cache.temporal[batch.mamba_track_indices]
|
||||
|
||||
torch.testing.assert_close(
|
||||
flashinfer_output, triton_output, atol=3e-2, rtol=3e-2
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
flashinfer_tracked, triton_tracked, atol=3e-2, rtol=3e-2
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -4,11 +4,18 @@ from unittest.mock import MagicMock, patch, sentinel
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
||||
MambaAttnBackendBase,
|
||||
)
|
||||
from sglang.srt.layers.attention.linear import gdn_backend
|
||||
from sglang.srt.layers.attention.linear.gdn_backend import (
|
||||
GDNAttnBackend,
|
||||
GDNKernelDispatcher,
|
||||
maybe_set_default_flashinfer_gdn_prefill,
|
||||
)
|
||||
from sglang.srt.layers.attention.linear.kernels.gdn_flashinfer import (
|
||||
maybe_build_flashinfer_checkpoint_plan,
|
||||
)
|
||||
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
|
||||
@@ -32,16 +39,13 @@ def make_runner(
|
||||
enable_dynamic_chunking=False,
|
||||
chunked_prefill_size=8192,
|
||||
)
|
||||
for name, value in arg_overrides.items():
|
||||
setattr(args, name, value)
|
||||
|
||||
# The policy routes its load-time default through the audited mutation entry
|
||||
# (server_args.override); mirror that on the stub so the write lands.
|
||||
def _override(source, **fields):
|
||||
for _field, _value in fields.items():
|
||||
setattr(args, _field, _value)
|
||||
|
||||
args.override = _override
|
||||
args.override = MagicMock(
|
||||
side_effect=lambda _source, **fields: vars(args).update(fields)
|
||||
)
|
||||
for name, value in arg_overrides.items():
|
||||
setattr(args, name, value)
|
||||
|
||||
return SimpleNamespace(
|
||||
server_args=args,
|
||||
@@ -87,14 +91,21 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
|
||||
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",
|
||||
)
|
||||
runner = make_runner()
|
||||
self.assertEqual(self.apply_policy(runner), "flashinfer")
|
||||
runner.server_args.override.assert_called_once_with(
|
||||
"gdn_backend.sm100_flashinfer_default",
|
||||
linear_attn_prefill_backend="flashinfer",
|
||||
)
|
||||
|
||||
def test_selects_flashinfer_for_radix_cache_strategies(self):
|
||||
for strategy in ("no_buffer", "extra_buffer", "extra_buffer_lazy"):
|
||||
with self.subTest(strategy=strategy):
|
||||
runner = make_runner(
|
||||
uses_mamba_radix_cache=True,
|
||||
mamba_radix_cache_strategy=strategy,
|
||||
)
|
||||
self.assertEqual(self.apply_policy(runner), "flashinfer")
|
||||
|
||||
def test_preserves_explicit_prefill_override(self):
|
||||
for backend in ("triton", "flashinfer", "cutedsl"):
|
||||
@@ -128,20 +139,6 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
|
||||
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}),
|
||||
@@ -151,6 +148,53 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
|
||||
with self.subTest(name=name):
|
||||
self.assertIsNone(self.apply_policy(make_runner(**runner_args)))
|
||||
|
||||
def test_builds_compact_checkpoint_plan_for_packed_sequences(self):
|
||||
forward_batch = SimpleNamespace(
|
||||
extend_seq_lens=torch.tensor([63, 64, 65, 127, 128, 129]),
|
||||
mamba_track_mask=torch.tensor([False, True, True, True, True, True]),
|
||||
# 65 on the 128-token sequence represents an interior S64
|
||||
# boundary encoded as S64 + 1 by the scheduler.
|
||||
mamba_track_seqlens=torch.tensor([63, 64, 65, 127, 65, 129]),
|
||||
extend_prefix_lens=torch.zeros(6, dtype=torch.int64),
|
||||
)
|
||||
metadata = SimpleNamespace(
|
||||
track_ssm_h_src=torch.empty(4),
|
||||
track_ssm_h_dst=torch.empty(4),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.layers.attention.linear.kernels.gdn_flashinfer."
|
||||
"get_server_args",
|
||||
return_value=SimpleNamespace(mamba_cache_chunk_size=64),
|
||||
):
|
||||
maybe_build_flashinfer_checkpoint_plan(forward_batch, metadata, "cpu")
|
||||
|
||||
torch.testing.assert_close(
|
||||
metadata.state_checkpoint_cu_starts,
|
||||
torch.tensor([0, 0, 1, 2, 3, 5, 7]),
|
||||
)
|
||||
torch.testing.assert_close(metadata.track_ssm_h_src, torch.tensor([1, 2, 3, 6]))
|
||||
self.assertEqual(metadata.num_state_checkpoints, 7)
|
||||
self.assertEqual(metadata.state_checkpoint_every_n_tokens, 64)
|
||||
|
||||
def test_decode_tracking_without_h_source_skips_checkpoint_plan(self):
|
||||
backend = object.__new__(GDNAttnBackend)
|
||||
backend.device = "cpu"
|
||||
backend.kernel_dispatcher = SimpleNamespace(extend_uses_state_checkpoints=True)
|
||||
metadata = SimpleNamespace(has_mamba_track_mask=True, track_ssm_h_src=None)
|
||||
forward_batch = SimpleNamespace(
|
||||
mamba_track_mask=torch.tensor([True]),
|
||||
mamba_track_indices=torch.tensor([7]),
|
||||
)
|
||||
|
||||
def init_base(instance, _forward_batch):
|
||||
instance.forward_metadata = metadata
|
||||
|
||||
with patch.object(MambaAttnBackendBase, "init_forward_metadata", init_base):
|
||||
backend.init_forward_metadata(forward_batch)
|
||||
|
||||
torch.testing.assert_close(metadata.conv_states_mask_indices, torch.tensor([7]))
|
||||
|
||||
def test_tree_verify_uses_triton_kernel(self):
|
||||
flashinfer_kernel = MagicMock(supports_target_verify=True)
|
||||
with (
|
||||
|
||||
Reference in New Issue
Block a user