[Refactor] Generalize DeepSeek V4 compressed pool management (#38954)
This commit is contained in:
@@ -45,6 +45,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
||||
):
|
||||
metadata = PagedIndexerMetadata(
|
||||
page_size=256,
|
||||
compressed_page_size=64,
|
||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||
use_topk_v2=False,
|
||||
@@ -59,22 +60,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
||||
self.assertEqual(args[1:], (64, 1))
|
||||
jit_metadata.assert_not_called()
|
||||
|
||||
def test_sm120_fp8_torch_fallback_keeps_metadata_none(self):
|
||||
with (
|
||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
|
||||
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
||||
envs.SGLANG_OPT_USE_TOPK_V2.override(False),
|
||||
):
|
||||
metadata = PagedIndexerMetadata(
|
||||
page_size=256,
|
||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||
use_topk_v2=False,
|
||||
)
|
||||
|
||||
self.assertIsNone(metadata.deep_gemm_metadata)
|
||||
|
||||
def test_topk_v2_ineligible_backend_skips_plan(self):
|
||||
def test_torch_fallback_skips_deep_gemm_and_ineligible_topk_plan(self):
|
||||
with (
|
||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
|
||||
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
||||
@@ -83,14 +69,54 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
||||
):
|
||||
metadata = PagedIndexerMetadata(
|
||||
page_size=256,
|
||||
compressed_page_size=64,
|
||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||
use_topk_v2=False,
|
||||
)
|
||||
|
||||
self.assertIsNone(metadata.deep_gemm_metadata)
|
||||
plan_topk_v2.assert_not_called()
|
||||
self.assertEqual(metadata.topk_metadata.numel(), 0)
|
||||
|
||||
def test_physical_page_size_controls_metadata_and_replay(self):
|
||||
planner = MagicMock(return_value=torch.zeros((1, 2), dtype=torch.int32))
|
||||
deep_gemm = SimpleNamespace(
|
||||
get_num_sms=MagicMock(return_value=1),
|
||||
get_paged_mqa_logits_metadata=planner,
|
||||
)
|
||||
with patch.dict(sys.modules, {"deep_gemm": deep_gemm}):
|
||||
metadata = [
|
||||
PagedIndexerMetadata(
|
||||
page_size=256,
|
||||
compressed_page_size=page_size,
|
||||
page_table=torch.zeros((1, 3), dtype=torch.int32),
|
||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||
use_topk_v2=False,
|
||||
force_deep_gemm_metadata=True,
|
||||
)
|
||||
for page_size in (64, 32, 32)
|
||||
]
|
||||
|
||||
self.assertEqual(
|
||||
[call.args[1] for call in planner.call_args_list], [64, 32, 32]
|
||||
)
|
||||
self.assertEqual([m.max_compressed_seq_len for m in metadata], [192, 96, 96])
|
||||
self.assertEqual([m.max_seq_len for m in metadata], [768, 768, 768])
|
||||
with self.assertRaisesRegex(AssertionError, "compressed_page_size"):
|
||||
metadata[0].copy_(metadata[1])
|
||||
|
||||
destination, source = metadata[1:]
|
||||
source.page_table.fill_(7)
|
||||
source.compressed_seq_lens.fill_(17)
|
||||
page_table_ptr = destination.page_table.data_ptr()
|
||||
destination.copy_(source)
|
||||
self.assertEqual(destination.page_table.data_ptr(), page_table_ptr)
|
||||
torch.testing.assert_close(destination.page_table, source.page_table)
|
||||
torch.testing.assert_close(
|
||||
destination.compressed_seq_lens, source.compressed_seq_lens
|
||||
)
|
||||
|
||||
|
||||
class TestDSV4FlashInferTopK(CustomTestCase):
|
||||
def test_compact_page_transform_respects_fuse_topk(self):
|
||||
|
||||
Reference in New Issue
Block a user