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:
co-authored by
hjzhang
hzh0425
parent
2f34dbe372
commit
c1b5c7e499
@@ -530,6 +530,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.qk_rope_head_dim = qk_rope_head_dim
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
self.indexer_head_dim = indexer_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)
|
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)
|
c128_layer_num = sum(1 for r in stage_ratios if r == 128)
|
||||||
c4_page_size = page_size // 4
|
c4_page_size = page_size // 4
|
||||||
@@ -572,7 +573,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
size=swa_size,
|
size=swa_size,
|
||||||
page_size=swa_page_size,
|
page_size=swa_page_size,
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
layer_num=layer_num,
|
layer_num=stage_layer_num,
|
||||||
device=device,
|
device=device,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
global_page_size=swa_page_size,
|
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."""
|
"""Convert absolute model layer_id to SWA-pool-local (PP-stage-local) index."""
|
||||||
return layer_id - self._stage_start
|
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:
|
def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
self.wait_layer_transfer(layer_id)
|
||||||
return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(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,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
# TODO(hzh0425): Support PP for deepseek v4 with hicache
|
|
||||||
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
|
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
|
||||||
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
||||||
swa_layer_mapping = {
|
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_num:
|
||||||
layer_id: layer_id for layer_id in range(len(kvcache.swa_kv_pool.kv_buffer))
|
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 = {}
|
c4_layer_mapping = {}
|
||||||
c128_layer_mapping = {}
|
c128_layer_mapping = {}
|
||||||
|
c4_state_local_layers = []
|
||||||
c4_state_global_layers = []
|
c4_state_global_layers = []
|
||||||
c128_state_global_layers = []
|
for local_layer_id, layer_item in enumerate(
|
||||||
for layer_id, layer_item in enumerate(
|
|
||||||
kvcache.layer_mapping[kvcache.start_layer : kvcache.end_layer]
|
kvcache.layer_mapping[kvcache.start_layer : kvcache.end_layer]
|
||||||
):
|
):
|
||||||
|
global_layer_id = kvcache.start_layer + local_layer_id
|
||||||
if layer_item.compress_ratio == 4:
|
if layer_item.compress_ratio == 4:
|
||||||
c4_layer_mapping[layer_id] = layer_item.compress_layer_id
|
c4_layer_mapping[local_layer_id] = layer_item.compress_layer_id
|
||||||
c4_state_global_layers.append(layer_id)
|
c4_state_local_layers.append(local_layer_id)
|
||||||
|
c4_state_global_layers.append(global_layer_id)
|
||||||
elif layer_item.compress_ratio == 128:
|
elif layer_item.compress_ratio == 128:
|
||||||
c128_layer_mapping[layer_id] = layer_item.compress_layer_id
|
c128_layer_mapping[local_layer_id] = layer_item.compress_layer_id
|
||||||
c128_state_global_layers.append(layer_id)
|
|
||||||
|
|
||||||
c4_state_mapping = {
|
c4_state_mapping = {
|
||||||
layer_id: local_id for local_id, layer_id in enumerate(c4_state_global_layers)
|
layer_id: local_id for local_id, layer_id in enumerate(c4_state_local_layers)
|
||||||
}
|
|
||||||
c128_state_mapping = {
|
|
||||||
layer_id: local_id for local_id, layer_id in enumerate(c128_state_global_layers)
|
|
||||||
}
|
}
|
||||||
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
|
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
|
||||||
params=params,
|
params=params,
|
||||||
|
|||||||
@@ -698,7 +698,7 @@ class MQALayer(nn.Module):
|
|||||||
|
|
||||||
token_to_kv_pool = get_token_to_kv_pool()
|
token_to_kv_pool = get_token_to_kv_pool()
|
||||||
swa_loc = attn_backend.get_swa_out_cache_loc(forward_batch)
|
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
|
swa_page_size = token_to_kv_pool.swa_kv_pool.page_size
|
||||||
|
|
||||||
q = fused_qk_norm_rope_swa_store(
|
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_loc = attn_backend.get_unified_swa_loc(forward_batch)
|
||||||
swa_page_size, bf16_store = 1, True
|
swa_page_size, bf16_store = 1, True
|
||||||
else:
|
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_loc = attn_backend.get_swa_out_cache_loc(forward_batch)
|
||||||
swa_page_size, bf16_store = (
|
swa_page_size, bf16_store = (
|
||||||
token_to_kv_pool.swa_kv_pool.page_size,
|
token_to_kv_pool.swa_kv_pool.page_size,
|
||||||
|
|||||||
+40
-29
@@ -34,6 +34,8 @@ def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label):
|
|||||||
class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||||
"""DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache."""
|
"""DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache."""
|
||||||
|
|
||||||
|
tp_size = 4
|
||||||
|
pp_size = 1
|
||||||
hicache_io_backend = "direct"
|
hicache_io_backend = "direct"
|
||||||
hicache_mem_layout = "page_first_direct"
|
hicache_mem_layout = "page_first_direct"
|
||||||
max_running_requests = 4
|
max_running_requests = 4
|
||||||
@@ -49,6 +51,43 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas
|
|||||||
def test_multiturn_logprobs_match(self):
|
def test_multiturn_logprobs_match(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _server_args(cls):
|
||||||
|
args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
str(cls.tp_size),
|
||||||
|
]
|
||||||
|
if cls.pp_size != 1:
|
||||||
|
args += ["--pp-size", str(cls.pp_size)]
|
||||||
|
args += [
|
||||||
|
"--attention-backend",
|
||||||
|
"compressed",
|
||||||
|
"--page-size",
|
||||||
|
"256",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"8192",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.9",
|
||||||
|
"--disable-shared-experts-fusion",
|
||||||
|
"--enable-hierarchical-cache",
|
||||||
|
"--hicache-ratio",
|
||||||
|
"4",
|
||||||
|
"--hicache-write-policy",
|
||||||
|
"write_through",
|
||||||
|
"--hicache-io-backend",
|
||||||
|
cls.hicache_io_backend,
|
||||||
|
"--hicache-mem-layout",
|
||||||
|
cls.hicache_mem_layout,
|
||||||
|
"--swa-full-tokens-ratio",
|
||||||
|
"0.25",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"20000",
|
||||||
|
"--max-running-requests",
|
||||||
|
str(cls.max_running_requests),
|
||||||
|
]
|
||||||
|
return args
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DSV4_FLASH_MODEL
|
cls.model = DSV4_FLASH_MODEL
|
||||||
@@ -57,35 +96,7 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas
|
|||||||
cls.model,
|
cls.model,
|
||||||
cls.base_url,
|
cls.base_url,
|
||||||
timeout=DSV4_FLASH_LAUNCH_TIMEOUT,
|
timeout=DSV4_FLASH_LAUNCH_TIMEOUT,
|
||||||
other_args=[
|
other_args=cls._server_args(),
|
||||||
"--trust-remote-code",
|
|
||||||
"--tp-size",
|
|
||||||
"4",
|
|
||||||
"--attention-backend",
|
|
||||||
"compressed",
|
|
||||||
"--page-size",
|
|
||||||
"256",
|
|
||||||
"--chunked-prefill-size",
|
|
||||||
"8192",
|
|
||||||
"--mem-fraction-static",
|
|
||||||
"0.9",
|
|
||||||
"--disable-shared-experts-fusion",
|
|
||||||
"--enable-hierarchical-cache",
|
|
||||||
"--hicache-ratio",
|
|
||||||
"4",
|
|
||||||
"--hicache-write-policy",
|
|
||||||
"write_through",
|
|
||||||
"--hicache-io-backend",
|
|
||||||
cls.hicache_io_backend,
|
|
||||||
"--hicache-mem-layout",
|
|
||||||
cls.hicache_mem_layout,
|
|
||||||
"--swa-full-tokens-ratio",
|
|
||||||
"0.25",
|
|
||||||
"--max-total-tokens",
|
|
||||||
"20000",
|
|
||||||
"--max-running-requests",
|
|
||||||
str(cls.max_running_requests),
|
|
||||||
],
|
|
||||||
env={
|
env={
|
||||||
"SGLANG_DSV4_FP4_EXPERTS": "0",
|
"SGLANG_DSV4_FP4_EXPERTS": "0",
|
||||||
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
|
"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()
|
||||||
Reference in New Issue
Block a user