Skip to content

Commit f926eea

Browse files
authored
[BugFix][KVTransfer] Fix Mooncake split metadata group ids (vllm-project#11439)
## What this PR does / why we need it? This PR fixes Mooncake KV split metadata generation for hybrid PD disaggregation scenarios where Mooncake transfer groups do not have a one-to-one mapping with KV cache manager groups. PR vllm-project#10991 split Mooncake transfer groups by KV spec and preserved the original KV cache manager group in `kv_cache_group_id`. PR vllm-project#10590 later moved kernel block expansion into `_get_kv_split_metadata`, but that path still indexed `local_block_ids` and `remote_block_ids` with the transfer group id. For hybrid models such as Kimi, this can make a Mamba or split attention transfer group index past the original KV cache group block-id tuple. The fix maps each transfer group back through `kv_cache_group_id` before reading or writing split metadata block ids, covering both no-CP and PCP/DCP paths. ## Does this PR introduce _any_ user-facing change? No user-facing API change. It fixes internal Mooncake KV transfer metadata generation for hybrid PD scenarios. ## How was this patch tested? CI and Unit tests. Unit tests: ```bash uv run pytest -q tests/ut/kv_offload/test_mooncake_connector.py ``` Result: 93 passed, 22 subtests passed. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@1f486d9 --------- Signed-off-by: zhuyixiang <zhuyixiang2014@163.com>
1 parent c120dad commit f926eea

2 files changed

Lines changed: 93 additions & 7 deletions

File tree

tests/ut/kv_offload/test_mooncake_connector.py

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2647,6 +2647,86 @@ def test_pd_disaggregated_hybrid_remote_pcp_splits_attention_and_final_mamba_sta
26472647
expected_finishes={0: 2, 1: 1},
26482648
)
26492649

2650+
def test_hybrid_no_cp_uses_kv_cache_group_ids_for_split_transfer_groups(self):
2651+
with patch.object(
2652+
self.vllm_config.kv_transfer_config,
2653+
"get_from_extra_config",
2654+
side_effect=lambda k, d=None: {
2655+
"prefill": {"tp_size": 4, "dp_size": 1, "pp_size": 1},
2656+
"decode": {"tp_size": 2, "dp_size": 1, "pp_size": 1},
2657+
}.get(k, d),
2658+
):
2659+
self.vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
2660+
self.vllm_config.model_config.is_deepseek_mla = False
2661+
self.vllm_config.model_config.hf_text_config.num_key_value_heads = 8
2662+
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
2663+
2664+
worker._is_hma_required = True
2665+
worker.use_mla = False
2666+
worker.use_sparse = False
2667+
worker.num_key_value_heads = 8
2668+
worker.tp_size = 2
2669+
worker.tp_rank = 0
2670+
worker.pcp_size = 1
2671+
worker.dcp_size = 1
2672+
worker.pcp_rank = 0
2673+
worker.dcp_rank = 0
2674+
worker._decode_tp_size = 2
2675+
worker._prefill_tp_size = 4
2676+
worker._prefill_pp_size = 1
2677+
worker.side_channel_port = 5000
2678+
worker.handshake_port = worker.side_channel_port + worker.tp_rank
2679+
worker.local_remote_block_port_mapping = {}
2680+
worker.remote_port_send_num = {}
2681+
worker.block_size_scale = [[1], [1], [1]]
2682+
worker.kv_group2layeridx = {
2683+
0: (
2684+
{
2685+
"kv_cache_spec_type": "FullAttentionSpec",
2686+
"kv_cache_group_id": 0,
2687+
"kv_cache_spec": {"num_kv_heads": 1},
2688+
},
2689+
[0],
2690+
),
2691+
1: (
2692+
{
2693+
"kv_cache_spec_type": "FullAttentionSpec",
2694+
"kv_cache_group_id": 0,
2695+
"kv_cache_spec": {"num_kv_heads": 8},
2696+
},
2697+
[1],
2698+
),
2699+
2: (
2700+
{
2701+
"kv_cache_spec_type": "MambaSpec",
2702+
"kv_cache_group_id": 1,
2703+
},
2704+
[2],
2705+
),
2706+
}
2707+
2708+
meta = types.SimpleNamespace(
2709+
remote_pcp_size=1,
2710+
remote_dcp_size=1,
2711+
remote_ptp_size=4,
2712+
remote_port=31000,
2713+
remote_block_ids=([50, 51, 52, 53], [60, 61, 62, 63]),
2714+
local_block_ids=([70, 71, 72, 73], [80, 81, 82, 83]),
2715+
num_external_tokens=4 * worker.block_size,
2716+
num_prompt_blocks=4,
2717+
num_computed_tokens=0,
2718+
remote_engine_id="remote_hybrid_split_transfer_groups",
2719+
remote_host="localhost",
2720+
remote_multi_nodes_meta_mapping={},
2721+
remote_block_size=16,
2722+
)
2723+
2724+
ports, local_ids, remote_ids = worker._get_kv_split_metadata("req_hybrid_split", cast(ReqMeta, meta))
2725+
2726+
self.assertEqual(len(ports), 1)
2727+
self.assertEqual(local_ids, [([70, 71, 72, 73], [80, 81, 82, 83])])
2728+
self.assertEqual(remote_ids, [([50, 51, 52, 53], [60, 61, 62, 63])])
2729+
26502730
def test_get_tp_num_need_pulls(self):
26512731
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
26522732
worker.num_key_value_heads = 8

vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py

Lines changed: 13 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2434,6 +2434,10 @@ def _group_compress_ratio(group_spec):
24342434
break
24352435
return compress_ratio
24362436

2437+
@staticmethod
2438+
def _get_kv_cache_group_id(group_idx: int, group_spec: dict[str, Any]) -> int:
2439+
return group_spec.get("kv_cache_group_id", group_idx)
2440+
24372441
def _get_kernel_block_ids(self, layer_indices, meta, group_idx, group_spec):
24382442
"""No-CP per-group block ids at kernel granularity: (local, remote).
24392443
@@ -2442,8 +2446,9 @@ def _get_kernel_block_ids(self, layer_indices, meta, group_idx, group_spec):
24422446
kernels (already on D, located via num_computed_tokens), and trims both lists
24432447
to the shorter one so remote/local stay aligned.
24442448
"""
2449+
kv_cache_group_id = self._get_kv_cache_group_id(group_idx, group_spec)
24452450
if group_spec["kv_cache_spec_type"] == "MambaSpec":
2446-
return list(meta.local_block_ids[group_idx]), list(meta.remote_block_ids[group_idx])
2451+
return list(meta.local_block_ids[kv_cache_group_id]), list(meta.remote_block_ids[kv_cache_group_id])
24472452

24482453
remote_block_size = meta.remote_block_size or self.block_size
24492454

@@ -2455,8 +2460,8 @@ def _get_kernel_block_ids(self, layer_indices, meta, group_idx, group_spec):
24552460
)
24562461

24572462
remote_scale = remote_block_size // kernel_size
2458-
kernel_local = self._expand_block_ids(list(meta.local_block_ids[group_idx]), local_scale)
2459-
kernel_remote = self._expand_block_ids(list(meta.remote_block_ids[group_idx]), remote_scale)
2463+
kernel_local = self._expand_block_ids(list(meta.local_block_ids[kv_cache_group_id]), local_scale)
2464+
kernel_remote = self._expand_block_ids(list(meta.remote_block_ids[kv_cache_group_id]), remote_scale)
24602465
# Skip prefix-cached remote kernels (D-side already holds them). The token size of one
24612466
# remote kernel is kernel_size * compress_ratio, so the number to skip is
24622467
# num_computed_tokens // (kernel_size * compress_ratio).
@@ -2551,14 +2556,15 @@ def _get_kv_split_metadata(
25512556
remote_handshake_port_list = [[x + meta.remote_port for x in chosen_rank_list]]
25522557
# No CP: expand logical blocks into kernel blocks here so the transfer
25532558
# stage consumes kernel-level ids directly (chunk_starts no longer needed).
2554-
local_block_ids: list = []
2555-
remote_block_ids: list = []
2559+
local_block_ids: list[list[int]] = [[] for _ in meta.local_block_ids]
2560+
remote_block_ids: list[list[int]] = [[] for _ in meta.remote_block_ids]
25562561
for group_idx, (group_spec, layer_indices) in self.kv_group2layeridx.items():
25572562
local_kernel_block_ids, remote_kernel_block_ids = self._get_kernel_block_ids(
25582563
layer_indices, meta, group_idx, group_spec
25592564
)
2560-
local_block_ids.append(local_kernel_block_ids)
2561-
remote_block_ids.append(remote_kernel_block_ids)
2565+
kv_cache_group_id = self._get_kv_cache_group_id(group_idx, group_spec)
2566+
local_block_ids[kv_cache_group_id] = local_kernel_block_ids
2567+
remote_block_ids[kv_cache_group_id] = remote_kernel_block_ids
25622568
local_block_ids_list = [tuple(local_block_ids) for _ in remote_handshake_port_list]
25632569
remote_block_ids_list = [tuple(remote_block_ids) for _ in remote_handshake_port_list]
25642570
return (

0 commit comments

Comments
 (0)