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:
@@ -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()
|
||||
Reference in New Issue
Block a user