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,