[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):
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
import unittest
|
||||
from itertools import product
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
|
||||
DeepSeekV4SingleKVPool,
|
||||
DeepSeekV4TokenToKVPool,
|
||||
_CompressedPoolConfig,
|
||||
)
|
||||
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 TestDSV4CompressedPools(CustomTestCase):
|
||||
def test_pp_mapping_and_pd_buffer_order(self):
|
||||
for unified, stage_ratios in product(
|
||||
(False, True), ([4, 0, 128, 4], [128], [0])
|
||||
):
|
||||
with self.subTest(unified=unified, stage_ratios=stage_ratios):
|
||||
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||
pool._unified_kv = unified
|
||||
pool.uniform_fp8 = False
|
||||
pool.compressed_pool_configs = {
|
||||
4: _CompressedPoolConfig(
|
||||
256, 64, torch.bfloat16, indexer_size=1024
|
||||
),
|
||||
128: _CompressedPoolConfig(512, 8, torch.float32),
|
||||
}
|
||||
pool.indexer_head_dim = 128
|
||||
pool.page_size = 256
|
||||
pool.compression_ratios = [128] + stage_ratios + [4]
|
||||
pool._stage_start = 1
|
||||
pool._stage_end = 1 + len(stage_ratios)
|
||||
|
||||
def kv_factory(**kwargs):
|
||||
return SimpleNamespace(
|
||||
kv_buffer=[
|
||||
torch.empty((3, kwargs["page_size"]), dtype=torch.uint8)
|
||||
for _ in range(kwargs["layer_num"])
|
||||
]
|
||||
)
|
||||
|
||||
# Separate payload/scale buffers exercise the FP4 transfer contract.
|
||||
indexer_buffers = [
|
||||
torch.empty((3, width), dtype=torch.uint8)
|
||||
for _ in range(stage_ratios.count(4))
|
||||
for width in (32, 2)
|
||||
]
|
||||
indexer = SimpleNamespace(
|
||||
contiguous_page_row_buffers=lambda: indexer_buffers
|
||||
)
|
||||
with (
|
||||
patch.object(pool, "_make_kv_pool", side_effect=kv_factory),
|
||||
patch.object(pool, "_make_indexer_pool", return_value=indexer),
|
||||
):
|
||||
pool._init_compressed_pools(
|
||||
stage_ratios=stage_ratios,
|
||||
page_size=256,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device="cpu",
|
||||
enable_memory_saver=False,
|
||||
enable_hisparse=False,
|
||||
kv_pool_cls=DeepSeekV4SingleKVPool,
|
||||
)
|
||||
pool._init_compressed_layer_mapping()
|
||||
self.assertIsNone(pool.layer_mapping[0])
|
||||
self.assertIsNone(pool.layer_mapping[-1])
|
||||
for local_id, ratio in enumerate(stage_ratios):
|
||||
item = pool.layer_mapping[local_id + 1]
|
||||
self.assertEqual(
|
||||
item.compress_layer_id, stage_ratios[:local_id].count(ratio)
|
||||
)
|
||||
self.assertIs(item.compress_kv_pool, pool.kv_pools.get(ratio))
|
||||
self.assertIs(pool.c4_kv_pool, pool.kv_pools[4])
|
||||
self.assertIs(pool.c128_kv_pool, pool.kv_pools[128])
|
||||
self.assertIs(pool.c4_indexer_kv_pool, pool.index_pools[4])
|
||||
|
||||
if unified:
|
||||
buffers = [
|
||||
torch.empty((9, 8), dtype=torch.uint8) for _ in stage_ratios
|
||||
]
|
||||
pool.unified_kv_pool = SimpleNamespace(
|
||||
swa_pages=2, kv_buffer=buffers
|
||||
)
|
||||
|
||||
def kv_entries(ratio):
|
||||
return [
|
||||
(buf.data_ptr() + 16, 56, 256 // ratio * 8)
|
||||
for buf, r in zip(buffers, stage_ratios)
|
||||
if r == ratio
|
||||
]
|
||||
else:
|
||||
|
||||
def kv_entries(ratio):
|
||||
return [
|
||||
(b.data_ptr(), b.nbytes, b[0].nbytes)
|
||||
for b in pool.kv_pools[ratio].kv_buffer
|
||||
]
|
||||
|
||||
indexer_entries = [
|
||||
(b.data_ptr(), b.nbytes, b[0].nbytes) for b in indexer_buffers
|
||||
]
|
||||
expected = kv_entries(4) + indexer_entries + kv_entries(128)
|
||||
actual = list(zip(*pool.get_contiguous_buf_infos()))
|
||||
self.assertEqual(actual, expected)
|
||||
|
||||
def test_shared_state_factory_preserves_layouts(self):
|
||||
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||
pool.compressed_pool_configs = {
|
||||
4: _CompressedPoolConfig(256, 64, torch.bfloat16, indexer_size=1024),
|
||||
128: _CompressedPoolConfig(512, 8, torch.float32),
|
||||
}
|
||||
pool.compression_ratios = [0, 4, 128]
|
||||
pool._stage_start, pool._stage_end = 0, 3
|
||||
pool.index_pools = {4: object()}
|
||||
pool.qk_nope_head_dim, pool.qk_rope_head_dim = 448, 64
|
||||
pool.indexer_head_dim = 128
|
||||
pool.device = "cpu"
|
||||
pool.swa_page_size = 128
|
||||
pool.online_mtp_max_draft_tokens = 3
|
||||
for online in (False, True):
|
||||
with (
|
||||
self.subTest(online=online),
|
||||
patch(
|
||||
"sglang.srt.mem_cache.deepseek_v4_memory_pool.ONLINE_C128", online
|
||||
),
|
||||
patch.object(
|
||||
pool,
|
||||
"get_ring_size",
|
||||
side_effect=lambda r: 8 if r == 4 else (1 if online else 128),
|
||||
),
|
||||
):
|
||||
pool._init_paged_compress_states(False)
|
||||
c4 = pool.compress_state_pools[1].kv_score_buffer.kv_score
|
||||
indexer = pool.indexer_compress_state_pools[1].kv_score_buffer.kv_score
|
||||
c128 = pool.compress_state_pools[2].kv_score_buffer.kv_score
|
||||
self.assertEqual(c4.shape, (76, 2048))
|
||||
self.assertEqual(indexer.shape, (76, 512))
|
||||
self.assertEqual(c128.shape, (40, 1536) if online else (256, 1024))
|
||||
self.assertEqual(c4.dtype, torch.bfloat16)
|
||||
self.assertEqual(indexer.dtype, torch.bfloat16)
|
||||
self.assertEqual(c128.dtype, torch.float32)
|
||||
self.assertNotEqual(c4.data_ptr(), indexer.data_ptr())
|
||||
self.assertIsNone(pool.compress_state_pools[0])
|
||||
self.assertIsNone(pool.indexer_compress_state_pools[2])
|
||||
|
||||
def test_indexer_access_uses_layer_ratio_and_waits_only_for_reads(self):
|
||||
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||
pool.kv_pools = {4: None, 128: None}
|
||||
pool.compression_ratios = [4, 128, 4]
|
||||
pool._stage_start, pool._stage_end = 0, 3
|
||||
pool._init_compressed_layer_mapping()
|
||||
indexer = MagicMock(page_size=32)
|
||||
pool.index_pools = {4: indexer}
|
||||
trace = MagicMock()
|
||||
trace.attach_mock(indexer, "indexer")
|
||||
with patch.object(pool, "wait_layer_transfer") as wait:
|
||||
trace.attach_mock(wait, "wait")
|
||||
pool.get_index_k_fp4_payload_buffer(2)
|
||||
pool.set_index_k_fp4(2, "loc", "cache")
|
||||
self.assertEqual(
|
||||
trace.mock_calls,
|
||||
[
|
||||
unittest.mock.call.wait(2),
|
||||
unittest.mock.call.indexer.get_index_k_fp4_payload_buffer(1),
|
||||
unittest.mock.call.indexer.set_index_fp4(1, "loc", "cache"),
|
||||
],
|
||||
)
|
||||
self.assertEqual(pool.get_index_k_page_size(4), 32)
|
||||
with self.assertRaisesRegex(AssertionError, "No indexer pool"):
|
||||
pool.get_index_k_page_size(128)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user