Fix Qwen3.5 GDN multi-item scoring (#33922)

Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
DAI0818
2026-09-09 22:30:39 -07:00
committed by GitHub
co-authored by Po-Han Huang
parent 6f481ad0e3
commit 03e4c06589
13 changed files with 842 additions and 27 deletions
@@ -13,6 +13,7 @@ from sglang.srt.layers.attention.linear.gdn_backend import (
GDNKernelDispatcher,
_validate_gdn_linear_attn_backends,
flashinfer_gdn_prefill_default,
validate_gdn_mis_backend,
)
from sglang.srt.layers.attention.linear.kernels.gdn_flashinfer import (
maybe_build_flashinfer_checkpoint_plan,
@@ -59,6 +60,7 @@ def make_runner(
mamba_radix_cache_strategy="no_buffer",
enable_dynamic_chunking=False,
chunked_prefill_size=8192,
enable_mis=False,
)
fields.update(arg_overrides)
args = _publish(testcase, **fields)
@@ -79,6 +81,22 @@ def make_runner(
class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase):
def test_mis_requires_triton_prefill_backend(self):
runner = make_runner(self, enable_mis=True)
with self.assertRaisesRegex(ValueError, "Triton linear-attention prefill"):
validate_gdn_mis_backend(LinearAttnKernelBackend.FLASHINFER)
def test_mis_rejects_page_major_layout(self):
make_runner(self, enable_mis=True, enable_page_major_kv_layout=True)
with self.assertRaisesRegex(ValueError, "page-major"):
validate_gdn_mis_backend(LinearAttnKernelBackend.TRITON)
def test_non_gdn_linear_backend_rejects_mis(self):
with self.assertRaisesRegex(ValueError, "does not support multi-item scoring"):
MambaAttnBackendBase.validate_mis_support(SimpleNamespace(enable_mis=True))
GDNAttnBackend.validate_mis_support(SimpleNamespace(enable_mis=True))
def apply_policy(
self,
runner,
@@ -232,6 +250,7 @@ class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase):
backend.kernel_dispatcher = SimpleNamespace(extend_uses_state_checkpoints=True)
metadata = SimpleNamespace(has_mamba_track_mask=True, track_ssm_h_src=None)
forward_batch = SimpleNamespace(
multi_item_delimiter_indices=None,
mamba_track_mask=torch.tensor([True]),
mamba_track_indices=torch.tensor([7]),
)
@@ -0,0 +1,99 @@
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.layers.attention.linear.gdn_backend import build_gdn_mis_metadata
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestGDNMISMetadata(CustomTestCase):
def test_allows_trailing_attention_padding(self):
forward_batch = SimpleNamespace(
input_ids=torch.empty(8, dtype=torch.int64),
extend_seq_lens_cpu=[5],
extend_prefix_lens_cpu=[0],
multi_item_delimiter_indices=[torch.tensor([1, 4], dtype=torch.int64)],
is_prefill_only=True,
)
metadata = build_gdn_mis_metadata(forward_batch)
torch.testing.assert_close(
torch.cat([metadata.query_token_indices, metadata.item_token_indices])
.sort()
.values,
torch.arange(5, dtype=torch.int64),
)
def test_mixed_batch_with_empty_query(self):
forward_batch = SimpleNamespace(
input_ids=torch.empty(16, dtype=torch.int64),
extend_seq_lens_cpu=[9, 7],
extend_prefix_lens_cpu=[0, 0],
multi_item_delimiter_indices=[
torch.tensor([0, 3, 8], dtype=torch.int64),
torch.tensor([4, 6], dtype=torch.int64),
],
is_prefill_only=True,
)
metadata = build_gdn_mis_metadata(forward_batch)
self.assertEqual(metadata.query_seq_lens_cpu, [4])
self.assertEqual(metadata.item_seq_lens_cpu, [3, 5, 1, 2, 1])
torch.testing.assert_close(
metadata.query_token_indices,
torch.tensor([9, 10, 11, 12], dtype=torch.int64),
)
torch.testing.assert_close(
metadata.query_cu_seqlens,
torch.tensor([0, 4], dtype=torch.int32),
)
torch.testing.assert_close(
metadata.query_request_indices,
torch.tensor([1], dtype=torch.int64),
)
torch.testing.assert_close(
metadata.item_token_indices,
torch.tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 13, 14, 15], dtype=torch.int64),
)
torch.testing.assert_close(
metadata.item_cu_seqlens,
torch.tensor([0, 3, 8, 9, 11, 12], dtype=torch.int32),
)
torch.testing.assert_close(
metadata.item_request_indices,
torch.tensor([0, 0, 0, 1, 1], dtype=torch.int64),
)
def test_rejects_sequence_lengths_beyond_input(self):
forward_batch = SimpleNamespace(
input_ids=torch.empty(4, dtype=torch.int64),
extend_seq_lens_cpu=[5],
extend_prefix_lens_cpu=[0],
multi_item_delimiter_indices=[torch.tensor([1, 4], dtype=torch.int64)],
is_prefill_only=True,
)
with self.assertRaisesRegex(ValueError, "exceed the input tokens"):
build_gdn_mis_metadata(forward_batch)
def test_rejects_missing_final_delimiter(self):
forward_batch = SimpleNamespace(
input_ids=torch.empty(5, dtype=torch.int64),
extend_seq_lens_cpu=[5],
extend_prefix_lens_cpu=[0],
multi_item_delimiter_indices=[torch.tensor([2, 3], dtype=torch.int64)],
is_prefill_only=True,
)
with self.assertRaisesRegex(ValueError, "final delimiter"):
build_gdn_mis_metadata(forward_batch)
if __name__ == "__main__":
unittest.main()