Fix DeepSeek V4 PP HiCache SWA allocation and layer mapping (#29106)

Co-authored-by: hjzhang <zhanghjzzz@qq.com>
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
hjzhang
2026-06-27 22:19:14 +08:00
committed by GitHub
co-authored by hjzhang hzh0425
parent 2f34dbe372
commit c1b5c7e499
5 changed files with 86 additions and 46 deletions
@@ -530,6 +530,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.qk_rope_head_dim = qk_rope_head_dim
self.indexer_head_dim = indexer_head_dim
stage_layer_num = len(stage_ratios)
c4_layer_num = sum(1 for r in stage_ratios if r == 4)
c128_layer_num = sum(1 for r in stage_ratios if r == 128)
c4_page_size = page_size // 4
@@ -572,7 +573,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
size=swa_size,
page_size=swa_page_size,
dtype=dtype,
layer_num=layer_num,
layer_num=stage_layer_num,
device=device,
enable_memory_saver=enable_memory_saver,
global_page_size=swa_page_size,
@@ -925,6 +926,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
"""Convert absolute model layer_id to SWA-pool-local (PP-stage-local) index."""
return layer_id - self._stage_start
def get_swa_raw_buffer(self, layer_id: int) -> torch.Tensor:
return self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)]
def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id)
return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id))
@@ -286,32 +286,33 @@ def build_deepseek_v4_hicache_stack(
storage_backend_extra_config: Optional[dict] = None,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
# TODO(hzh0425): Support PP for deepseek v4 with hicache
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
swa_layer_mapping = {
layer_id: layer_id for layer_id in range(len(kvcache.swa_kv_pool.kv_buffer))
}
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_num:
raise ValueError(
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
f"got {len(kvcache.swa_kv_pool.kv_buffer)} buffers for "
f"{transfer_layer_num} local layers"
)
swa_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
c4_layer_mapping = {}
c128_layer_mapping = {}
c4_state_local_layers = []
c4_state_global_layers = []
c128_state_global_layers = []
for layer_id, layer_item in enumerate(
for local_layer_id, layer_item in enumerate(
kvcache.layer_mapping[kvcache.start_layer : kvcache.end_layer]
):
global_layer_id = kvcache.start_layer + local_layer_id
if layer_item.compress_ratio == 4:
c4_layer_mapping[layer_id] = layer_item.compress_layer_id
c4_state_global_layers.append(layer_id)
c4_layer_mapping[local_layer_id] = layer_item.compress_layer_id
c4_state_local_layers.append(local_layer_id)
c4_state_global_layers.append(global_layer_id)
elif layer_item.compress_ratio == 128:
c128_layer_mapping[layer_id] = layer_item.compress_layer_id
c128_state_global_layers.append(layer_id)
c128_layer_mapping[local_layer_id] = layer_item.compress_layer_id
c4_state_mapping = {
layer_id: local_id for local_id, layer_id in enumerate(c4_state_global_layers)
}
c128_state_mapping = {
layer_id: local_id for local_id, layer_id in enumerate(c128_state_global_layers)
layer_id: local_id for local_id, layer_id in enumerate(c4_state_local_layers)
}
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
params=params,
+2 -2
View File
@@ -698,7 +698,7 @@ class MQALayer(nn.Module):
token_to_kv_pool = get_token_to_kv_pool()
swa_loc = attn_backend.get_swa_out_cache_loc(forward_batch)
swa_cache = token_to_kv_pool.swa_kv_pool.kv_buffer[self.layer_id]
swa_cache = token_to_kv_pool.get_swa_raw_buffer(self.layer_id)
swa_page_size = token_to_kv_pool.swa_kv_pool.page_size
q = fused_qk_norm_rope_swa_store(
@@ -799,7 +799,7 @@ class MQALayer(nn.Module):
swa_loc = attn_backend.get_unified_swa_loc(forward_batch)
swa_page_size, bf16_store = 1, True
else:
swa_cache = token_to_kv_pool.swa_kv_pool.kv_buffer[self.layer_id]
swa_cache = token_to_kv_pool.get_swa_raw_buffer(self.layer_id)
swa_loc = attn_backend.get_swa_out_cache_loc(forward_batch)
swa_page_size, bf16_store = (
token_to_kv_pool.swa_kv_pool.page_size,
@@ -34,6 +34,8 @@ def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label):
class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache."""
tp_size = 4
pp_size = 1
hicache_io_backend = "direct"
hicache_mem_layout = "page_first_direct"
max_running_requests = 4
@@ -50,17 +52,15 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas
pass
@classmethod
def setUpClass(cls):
cls.model = DSV4_FLASH_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DSV4_FLASH_LAUNCH_TIMEOUT,
other_args=[
def _server_args(cls):
args = [
"--trust-remote-code",
"--tp-size",
"4",
str(cls.tp_size),
]
if cls.pp_size != 1:
args += ["--pp-size", str(cls.pp_size)]
args += [
"--attention-backend",
"compressed",
"--page-size",
@@ -85,7 +85,18 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas
"20000",
"--max-running-requests",
str(cls.max_running_requests),
],
]
return args
@classmethod
def setUpClass(cls):
cls.model = DSV4_FLASH_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DSV4_FLASH_LAUNCH_TIMEOUT,
other_args=cls._server_args(),
env={
"SGLANG_DSV4_FP4_EXPERTS": "0",
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
@@ -0,0 +1,24 @@
import unittest
import test_unified_radix_cache_kl_dsv4 as dsv4_kl
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=900, stage="extra-b", runner_config="8-gpu-h200")
class TestUnifiedDeepSeekV4FlashHiCachePP4TP2(
dsv4_kl.TestUnifiedDeepSeekV4FlashHiCache
):
"""DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache under PP4 TP2."""
pp_size = 4
tp_size = 2
@unittest.skip("PP4TP2 coverage uses accuracy and cache-hit KL cases.")
def test_multiturn_logprobs_match(self):
pass
if __name__ == "__main__":
unittest.main()