[DeepSeek-V4] Enable non-paged indexer by default for large prefill chunks (#30140)
This commit is contained in:
@@ -842,7 +842,10 @@ class Envs:
|
|||||||
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
|
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
|
||||||
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
||||||
SGLANG_OPT_USE_AITER_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_JIT_INDEXER_METADATA = EnvBool(True)
|
||||||
SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False)
|
SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False)
|
||||||
SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False)
|
SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False)
|
||||||
|
|||||||
@@ -482,6 +482,8 @@ class C4IndexerBackendMixin:
|
|||||||
c4_seq_lens: torch.Tensor,
|
c4_seq_lens: torch.Tensor,
|
||||||
query_rows: int,
|
query_rows: int,
|
||||||
) -> Optional[NonPagedIndexerPlan]:
|
) -> Optional[NonPagedIndexerPlan]:
|
||||||
|
if query_rows < envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS.get():
|
||||||
|
return None
|
||||||
if not self._can_use_nonpaged_indexer(
|
if not self._can_use_nonpaged_indexer(
|
||||||
c4_indexer=c4_indexer,
|
c4_indexer=c4_indexer,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
|
|||||||
@@ -55,7 +55,10 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def test_eligibility_is_fail_closed(self):
|
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())
|
self.assertTrue(self._is_eligible())
|
||||||
for case in (
|
for case in (
|
||||||
{"enabled": False},
|
{"enabled": False},
|
||||||
@@ -98,6 +101,10 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
|
|||||||
query_rows=query_rows,
|
query_rows=query_rows,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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()
|
plan = build_plan()
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
(plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows),
|
(plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows),
|
||||||
@@ -109,6 +116,7 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
|
|||||||
|
|
||||||
metadata.nonpaged_plan = None
|
metadata.nonpaged_plan = None
|
||||||
batch.extend_seq_lens_cpu = [2, 2]
|
batch.extend_seq_lens_cpu = [2, 2]
|
||||||
|
with threshold.override(0):
|
||||||
self.assertIsNone(build_plan())
|
self.assertIsNone(build_plan())
|
||||||
|
|
||||||
def test_extreme_plan_metadata_is_bounded_and_fail_closed(self):
|
def test_extreme_plan_metadata_is_bounded_and_fail_closed(self):
|
||||||
@@ -140,6 +148,8 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
|
|||||||
query_rows=query_rows,
|
query_rows=query_rows,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
threshold = envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS
|
||||||
|
with threshold.override(query_rows):
|
||||||
plan = build_plan()
|
plan = build_plan()
|
||||||
self.assertEqual(plan.seq_len_sum, 125_000)
|
self.assertEqual(plan.seq_len_sum, 125_000)
|
||||||
self.assertEqual(plan.max_seq_len, 125_000)
|
self.assertEqual(plan.max_seq_len, 125_000)
|
||||||
@@ -151,8 +161,54 @@ class TestDSV4NonPagedIndexer(CustomTestCase):
|
|||||||
batch.extend_seq_lens_cpu = [2, 2]
|
batch.extend_seq_lens_cpu = [2, 2]
|
||||||
batch.extend_seq_lens = torch.tensor([2, 2], dtype=torch.int32)
|
batch.extend_seq_lens = torch.tensor([2, 2], dtype=torch.int32)
|
||||||
batch.extend_start_loc = torch.tensor([0, 2], dtype=torch.int32)
|
batch.extend_start_loc = torch.tensor([0, 2], dtype=torch.int32)
|
||||||
|
with threshold.override(query_rows):
|
||||||
self.assertIsNone(build_plan())
|
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):
|
def test_nonpaged_dispatch_uses_gathered_kv_contract(self):
|
||||||
query_rows = 4
|
query_rows = 4
|
||||||
plan = NonPagedIndexerPlan(
|
plan = NonPagedIndexerPlan(
|
||||||
|
|||||||
Reference in New Issue
Block a user