diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 2942f53da..edcc34041 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -842,7 +842,10 @@ class Envs: SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False) SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False) SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False) - SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(False) + SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(True) + # Per-rank local query rows (after DP-attention sharding when enabled), + # not request ISL. + SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS = EnvInt(8192) SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True) SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False) SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 575b26e8b..ed85e3da6 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -482,6 +482,8 @@ class C4IndexerBackendMixin: c4_seq_lens: torch.Tensor, query_rows: int, ) -> Optional[NonPagedIndexerPlan]: + if query_rows < envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS.get(): + return None if not self._can_use_nonpaged_indexer( c4_indexer=c4_indexer, forward_batch=forward_batch, diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index d70059c5d..fff3ad86f 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -55,7 +55,10 @@ class TestDSV4NonPagedIndexer(CustomTestCase): ) def test_eligibility_is_fail_closed(self): - self.assertIs(envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.default, False) + self.assertIs(envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.default, True) + self.assertEqual( + envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS.default, 8192 + ) self.assertTrue(self._is_eligible()) for case in ( {"enabled": False}, @@ -98,7 +101,11 @@ class TestDSV4NonPagedIndexer(CustomTestCase): query_rows=query_rows, ) - plan = build_plan() + threshold = envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS + with threshold.override(threshold.default): + self.assertIsNone(build_plan()) + with threshold.override(query_rows): + plan = build_plan() self.assertEqual( (plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows), (65, 128, query_rows), @@ -109,7 +116,8 @@ class TestDSV4NonPagedIndexer(CustomTestCase): metadata.nonpaged_plan = None batch.extend_seq_lens_cpu = [2, 2] - self.assertIsNone(build_plan()) + with threshold.override(0): + self.assertIsNone(build_plan()) def test_extreme_plan_metadata_is_bounded_and_fail_closed(self): backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True) @@ -140,7 +148,9 @@ class TestDSV4NonPagedIndexer(CustomTestCase): query_rows=query_rows, ) - plan = build_plan() + threshold = envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS + with threshold.override(query_rows): + plan = build_plan() self.assertEqual(plan.seq_len_sum, 125_000) self.assertEqual(plan.max_seq_len, 125_000) self.assertEqual(plan.max_seqlen_k, 125_056) @@ -151,7 +161,53 @@ class TestDSV4NonPagedIndexer(CustomTestCase): batch.extend_seq_lens_cpu = [2, 2] batch.extend_seq_lens = torch.tensor([2, 2], dtype=torch.int32) batch.extend_start_loc = torch.tensor([0, 2], dtype=torch.int32) - self.assertIsNone(build_plan()) + with threshold.override(query_rows): + self.assertIsNone(build_plan()) + + def test_query_threshold_boundary(self): + can_use_nonpaged_indexer = MagicMock(return_value=True) + backend = SimpleNamespace(_can_use_nonpaged_indexer=can_use_nonpaged_indexer) + c4_indexer = SimpleNamespace(use_fp4_indexer=False) + metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + + def build_plan(query_rows): + batch = SimpleNamespace( + seq_lens=torch.tensor([query_rows], dtype=torch.int32), + seq_lens_cpu=[query_rows], + extend_seq_lens_cpu=[query_rows], + extend_seq_lens=torch.tensor([query_rows], dtype=torch.int32), + extend_start_loc=torch.tensor([0], dtype=torch.int32), + extend_num_tokens=query_rows, + ) + c4_seq_lens = torch.div( + torch.arange(1, query_rows + 1, dtype=torch.int32), + 4, + rounding_mode="floor", + ).clamp_min_(1) + return C4IndexerBackendMixin._get_nonpaged_indexer_plan( + backend, + c4_indexer=c4_indexer, + forward_batch=batch, + indexer_metadata=metadata, + page_table=torch.zeros((query_rows, 1), dtype=torch.int32), + c4_seq_lens=c4_seq_lens, + query_rows=query_rows, + ) + + for query_rows, expected in ((8191, False), (8192, True), (8193, True)): + with self.subTest(query_rows=query_rows): + metadata.nonpaged_plan = None + can_use_nonpaged_indexer.reset_mock() + self.assertIs(build_plan(query_rows) is not None, expected) + if expected: + can_use_nonpaged_indexer.assert_called_once() + else: + can_use_nonpaged_indexer.assert_not_called() + + metadata.nonpaged_plan = None + threshold = envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS + with threshold.override(8193): + self.assertIsNone(build_plan(8192)) def test_nonpaged_dispatch_uses_gathered_kv_contract(self): query_rows = 4