Skip to content

Commit a7f9b9b

Browse files
[Feature] Fix Mooncake HMA KV spec group mapping (vllm-project#10991)
## What this PR does / why we need it? This PR fixes Mooncake KV transfer group mapping for HMA scenarios where the target model and Eagle draft model share one KV cache manager group but use different KV specs. In the MLA + QGA case, the main model has 1 KV head while the Eagle model has 8 KV heads. With P/D TP set to tp8/tp4, treating both specs as one transfer group causes incorrect TP pull mapping. This PR makes Mooncake transfer metadata distinguish KV specs by KV head count while preserving the mapping back to the original KV cache manager group. It also fixes P-side delayed-free rank calculation so the prefill worker computes the union of prefill ranks that may be pulled by any decode rank, instead of reusing D-side group pull logic. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@dc68bd8 --------- Signed-off-by: liziyu <liziyu16@huawei.com> Signed-off-by: wangxiaoteng <wangxiaoteng@huawei.com> Co-authored-by: wangxiaoteng <wangxiaoteng@huawei.com>
1 parent fd35fa5 commit a7f9b9b

2 files changed

Lines changed: 344 additions & 32 deletions

File tree

tests/ut/kv_offload/test_mooncake_connector.py

Lines changed: 140 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
import torch
1515
import zmq
1616
from vllm.utils.network_utils import make_zmq_path
17+
from vllm.v1.kv_cache_interface import FullAttentionSpec, UniformTypeKVCacheSpecs
1718

1819
fake_engine = types.ModuleType("mooncake.engine")
1920
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
@@ -295,6 +296,145 @@ def test_reformat_kv_cache_hybrid_linear_uses_cache_block_size(self):
295296
torch.testing.assert_close(reformatted_v_cache, expected)
296297

297298

299+
class TestMooncakeTransferGroups(unittest.TestCase):
300+
def test_build_kv_group2layeridx_splits_uniform_group_by_kv_heads(self):
301+
mla_spec = FullAttentionSpec(
302+
block_size=16,
303+
num_kv_heads=1,
304+
head_size=64,
305+
head_size_v=64,
306+
dtype=torch.float16,
307+
)
308+
qga_spec = FullAttentionSpec(
309+
block_size=16,
310+
num_kv_heads=8,
311+
head_size=64,
312+
head_size_v=64,
313+
dtype=torch.float16,
314+
)
315+
layer_specs = {
316+
"model.layers.0.self_attn": mla_spec,
317+
"model.layers.1.self_attn": qga_spec,
318+
}
319+
uniform_spec = UniformTypeKVCacheSpecs.from_specs(layer_specs)
320+
assert uniform_spec is not None
321+
322+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
323+
worker.vllm_config = MockVllmConfig()
324+
worker.total_layers = 32
325+
worker.kv_cache_config = MockKVCacheConfig(
326+
kv_cache_groups=[
327+
MockKVCacheGroup(
328+
layer_names=list(layer_specs),
329+
kv_cache_spec=uniform_spec,
330+
)
331+
]
332+
)
333+
334+
kv_group2layeridx = worker._build_kv_group2layeridx()
335+
336+
self.assertEqual(len(kv_group2layeridx), 2)
337+
self.assertEqual(kv_group2layeridx[0][0]["kv_cache_group_id"], 0)
338+
self.assertEqual(kv_group2layeridx[1][0]["kv_cache_group_id"], 0)
339+
self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[0][0]), 1)
340+
self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[1][0]), 8)
341+
342+
def test_hybrid_rank_pulls_use_transfer_group_kv_heads(self):
343+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
344+
worker.vllm_config = MockVllmConfig()
345+
worker.vllm_config.model_config.is_deepseek_mla = True
346+
worker.tp_rank = 0
347+
worker.tp_size = 4
348+
worker._decode_tp_size = 4
349+
worker._prefill_tp_size = 8
350+
worker._prefill_pp_size = 1
351+
worker.num_key_value_heads = 128
352+
worker.use_sparse = False
353+
worker.kv_group2layeridx = {
354+
0: (
355+
{
356+
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
357+
"kv_cache_group_id": 0,
358+
"kv_cache_spec": {"model.layers.0.self_attn": {"num_kv_heads": 1}},
359+
},
360+
[0],
361+
),
362+
1: (
363+
{
364+
"kv_cache_spec_type": "UniformTypeKVCacheSpecs",
365+
"kv_cache_group_id": 0,
366+
"kv_cache_spec": {"model.layers.1.self_attn": {"num_kv_heads": 8}},
367+
},
368+
[1],
369+
),
370+
}
371+
372+
_, rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls("req-1", prefill_tp_size=8)
373+
pulls = [pull for group_pulls in rank_group_pulls.values() for pull in group_pulls]
374+
mla_pulls = [pull for pull in pulls if pull.group_id == 0]
375+
qga_pulls = [pull for pull in pulls if pull.group_id == 1]
376+
377+
self.assertEqual(len(mla_pulls), 1)
378+
self.assertEqual(mla_pulls[0].num_group_pulls, 1)
379+
self.assertEqual(len(qga_pulls), 2)
380+
self.assertTrue(all(pull.num_group_pulls == 2 for pull in qga_pulls))
381+
382+
def test_hybrid_group_pulls_metadata_filters_groups_per_remote_card(self):
383+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
384+
worker.vllm_config = MockVllmConfig()
385+
worker.vllm_config.model_config.is_deepseek_mla = True
386+
worker._is_hma_required = True
387+
worker.tp_rank = 0
388+
worker.tp_size = 4
389+
worker._decode_tp_size = 4
390+
worker._prefill_tp_size = 8
391+
worker._prefill_pp_size = 1
392+
worker.num_key_value_heads = 128
393+
worker.use_sparse = False
394+
worker.kv_group2layeridx = {
395+
0: (
396+
{
397+
"kv_cache_spec_type": "FullAttentionSpec",
398+
"kv_cache_group_id": 0,
399+
"kv_cache_spec": {"num_kv_heads": 1},
400+
},
401+
[0],
402+
),
403+
1: (
404+
{
405+
"kv_cache_spec_type": "FullAttentionSpec",
406+
"kv_cache_group_id": 0,
407+
"kv_cache_spec": {"num_kv_heads": 8},
408+
},
409+
[1],
410+
),
411+
}
412+
req_id = "req-1"
413+
remote_base_port = 30000
414+
remote_handshake_port_list = [[remote_base_port + rank] for rank in range(8)]
415+
416+
_, expected_rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls(req_id, prefill_tp_size=8)
417+
group_pulls_list = worker._get_group_pulls_metadata(
418+
req_id,
419+
remote_handshake_port_list,
420+
prefill_tp_size=8,
421+
remote_base_port=remote_base_port,
422+
)
423+
424+
group_ids_by_rank = [
425+
[group_pull.group_id for group_pull in group_pulls_list[rank][0]]
426+
for rank in range(len(remote_handshake_port_list))
427+
]
428+
expected_group_ids_by_rank = [
429+
[group_pull.group_id for group_pull in expected_rank_group_pulls.get(rank, [])]
430+
for rank in range(len(remote_handshake_port_list))
431+
]
432+
433+
self.assertEqual(group_ids_by_rank, expected_group_ids_by_rank)
434+
self.assertTrue(any(set(group_ids) != {0, 1} for group_ids in group_ids_by_rank))
435+
self.assertFalse(all(set(group_ids) == {0, 1} for group_ids in group_ids_by_rank))
436+
437+
298438
class TestKVCacheRecvingThreadBasic(unittest.TestCase):
299439
def setUp(self):
300440
self.engine = MagicMock()

0 commit comments

Comments
 (0)