Skip to content

Commit c062a7e

Browse files
authored
[Feature][P/D] Support PD disaggregation for DCP with replicate-indexer for SFA (vllm-project#11696)
### What this PR does / why we need it? This PR adds P/D disaggregation support for SFA DCP replicate-indexer feature. vllm-project#11443 ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? By CI. - vLLM version: v0.24.0 - vLLM main: vllm-project/vllm@85c09e9 --------- Signed-off-by: nwpu-zxr <zhouxuerong2@huawei.com>
1 parent e881974 commit c062a7e

2 files changed

Lines changed: 451 additions & 49 deletions

File tree

tests/ut/kv_offload/test_mooncake_connector.py

Lines changed: 254 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1010,6 +1010,32 @@ def test_transfer_kv_cache_uses_block_stride_for_block_offsets(self, mock_get_me
10101010
self.assertEqual(call_args[3], [1024, 1024])
10111011
mock_get_meta.assert_not_called()
10121012

1013+
@patch.object(KVCacheRecvingThread, "_get_remote_metadata")
1014+
def test_transfer_replicated_indexer_when_regular_kv_shard_is_empty(self, mock_get_meta):
1015+
req = dict(self.test_req)
1016+
req["local_block_ids"] = [[]]
1017+
req["remote_block_ids"] = [[]]
1018+
req["local_block_ids_replicate_k"] = ([4, 5],)
1019+
req["remote_block_ids_replicate_k"] = ([7, 8],)
1020+
1021+
with patch("vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_connector.get_ascend_config") as mock_config:
1022+
mock_config.return_value.enable_kv_nz = False
1023+
self.thread.enable_sfa_dcp_replicated_indexer = True
1024+
self.thread.kv_caches_base_addr["local_engine"][5555] = [[0x1000, 0x2000]]
1025+
self.thread.kv_caches_base_addr["remote_engine"] = {6666: [[0x3000, 0x4000]]}
1026+
self.thread.block_size_scale = [[1, 2]]
1027+
self.thread.block_len_per_addr = [[1024, 2048]]
1028+
self.thread.block_stride_per_addr = [[1024, 2048]]
1029+
self.thread.remote_block_stride_per_addr["remote_engine"][6666] = [[4096, 8192]]
1030+
1031+
self.thread._transfer_kv_cache_all_groups(req)
1032+
1033+
call_args, _ = self.engine.batch_transfer_sync_read.call_args
1034+
self.assertEqual(call_args[1], [0x2000 + 4 * 2048])
1035+
self.assertEqual(call_args[2], [0x4000 + 7 * 8192])
1036+
self.assertEqual(call_args[3], [2 * 2048])
1037+
mock_get_meta.assert_not_called()
1038+
10131039
def test_append_mamba_transfer_meta_uses_block_stride_for_block_offsets(self):
10141040
src_list: list[int] = []
10151041
dst_list: list[int] = []
@@ -1296,6 +1322,7 @@ def test_add_new_req(self):
12961322
meta.add_new_req(
12971323
request_id="req1",
12981324
local_block_ids=[1, 2, 3],
1325+
local_full_block_ids=[0, 1, 2, 3],
12991326
num_external_tokens=48,
13001327
kv_transfer_params={
13011328
"remote_block_ids": [4, 5, 6],
@@ -1313,6 +1340,7 @@ def test_add_new_req(self):
13131340
req_meta = meta.requests["req1"]
13141341
self.assertIsInstance(req_meta, ReqMeta)
13151342
self.assertEqual(req_meta.local_block_ids, [1, 2, 3])
1343+
self.assertEqual(req_meta.local_full_block_ids, [0, 1, 2, 3])
13161344
self.assertEqual(req_meta.remote_block_ids, [4, 5, 6])
13171345
self.assertEqual(req_meta.remote_engine_id, "remote_engine")
13181346
self.assertEqual(req_meta.remote_host, "localhost")
@@ -1350,9 +1378,7 @@ def test_get_num_new_matched_tokens(self):
13501378

13511379
def test_build_connector_meta(self):
13521380
request = MockRequest("req1")
1353-
blocks_mock = MagicMock()
1354-
blocks_mock.get_unhashed_block_ids.return_value = [4, 5, 6]
1355-
self.scheduler._reqs_need_recv["req1"] = (request, [4, 5, 6], 48)
1381+
self.scheduler._reqs_need_recv["req1"] = (request, [4, 5, 6], [0, 4, 5, 6], 48)
13561382
request.kv_transfer_params = {
13571383
"remote_block_ids": [1, 2, 3],
13581384
"remote_engine_id": "remote",
@@ -1368,6 +1394,7 @@ def test_build_connector_meta(self):
13681394
self.assertIsInstance(meta, MooncakeConnectorMetadata)
13691395
self.assertEqual(len(meta.requests), 1)
13701396
self.assertEqual(meta.requests["req1"].local_block_ids, [4, 5, 6])
1397+
self.assertEqual(meta.requests["req1"].local_full_block_ids, [0, 4, 5, 6])
13711398
self.assertEqual(meta.requests["req1"].remote_block_ids, [1, 2, 3])
13721399
self.assertEqual(meta.requests["req1"].num_computed_tokens, 16)
13731400
self.assertEqual(len(self.scheduler._reqs_need_recv), 0)
@@ -1469,6 +1496,9 @@ def get_unhashed_block_ids(self):
14691496
def get_unhashed_block_ids_all_groups(self):
14701497
return ([4, 5, 6],)
14711498

1499+
def get_block_ids(self):
1500+
return ([1, 2, 4, 5, 6],)
1501+
14721502

14731503
class MockSchedulerOutput:
14741504
pass
@@ -1608,6 +1638,7 @@ def test_update_state_after_alloc_with_remote_prefill(self):
16081638
self.assertEqual(len(self.scheduler._reqs_need_recv), 1)
16091639
self.assertEqual(self.scheduler._reqs_need_recv["req1"][0], request)
16101640
self.assertEqual(self.scheduler._reqs_need_recv["req1"][1], ([4, 5, 6],))
1641+
self.assertEqual(self.scheduler._reqs_need_recv["req1"][2], ([1, 2, 4, 5, 6],))
16111642

16121643
def test_request_finished_no_remote_decode(self):
16131644
request = MockRequest("req1")
@@ -1654,6 +1685,21 @@ def test_get_transfer_block_ids_uses_compressed_prompt_len(self):
16541685

16551686
self.assertEqual(block_ids, ([30, 31],))
16561687

1688+
def test_get_transfer_block_ids_uses_cp_grouped_block_len(self):
1689+
self.scheduler.pcp_size = 1
1690+
self.scheduler.dcp_size = 4
1691+
self.scheduler.group_transfer_info = [
1692+
types.SimpleNamespace( # type: ignore[list-item]
1693+
tokens_per_block=16,
1694+
blocks_per_window=0,
1695+
is_state_group=False,
1696+
)
1697+
]
1698+
1699+
block_ids = self.scheduler._get_transfer_block_ids(([10, 11, 12, 13, 14],), prompt_len=65)
1700+
1701+
self.assertEqual(block_ids, ([10, 11],))
1702+
16571703
def test_get_transfer_block_ids_trims_sliding_window_mtp_blocks(self):
16581704
self.scheduler.group_transfer_info = [
16591705
types.SimpleNamespace( # type: ignore[list-item]
@@ -1726,6 +1772,28 @@ def test_request_finished_trims_mtp_blocks_in_params(self):
17261772
self.assertEqual(params["num_prompt_blocks"], 3)
17271773
self.assertIn("req_mtp", self.scheduler._reqs_need_send)
17281774

1775+
def test_request_finished_trims_cp_grouped_mtp_blocks_in_params(self):
1776+
self.scheduler.pcp_size = 1
1777+
self.scheduler.dcp_size = 4
1778+
self.scheduler.group_transfer_info = [
1779+
types.SimpleNamespace(
1780+
tokens_per_block=16,
1781+
blocks_per_window=0,
1782+
is_state_group=False,
1783+
)
1784+
]
1785+
request = self._make_remote_decode_request(prompt_len=65, request_id="req_cp_mtp")
1786+
1787+
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
1788+
1789+
self.assertTrue(delay_free)
1790+
self.assertIsNotNone(params)
1791+
assert params is not None
1792+
self.assertEqual(params["remote_block_ids"], ([10, 11],))
1793+
# num_prompt_blocks stays in no-CP units for worker-side CP distribution.
1794+
self.assertEqual(params["num_prompt_blocks"], 5)
1795+
self.assertIn("req_cp_mtp", self.scheduler._reqs_need_send)
1796+
17291797
def test_request_finished_clips_sliding_window_blocks_in_params(self):
17301798
self.scheduler.group_transfer_info = [
17311799
types.SimpleNamespace(
@@ -2914,6 +2982,189 @@ def test_get_tp_num_need_pulls(self):
29142982
tp_num_need_pulls = worker._get_tp_num_need_pulls(prefill_tp_size=None)
29152983
self.assertEqual(tp_num_need_pulls, 1)
29162984

2985+
def test_start_load_kv_puts_replicated_indexer_on_existing_transfer_port(self):
2986+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
2987+
worker.kv_send_thread = None
2988+
worker.kv_recv_thread = MagicMock()
2989+
worker._prefill_tp_size = 4
2990+
worker.remote_port_send_num = {"remote_engine": {31001: {"num": 1, "host": "localhost"}}}
2991+
worker._get_sfa_replicate_k_block_ids = MagicMock(return_value=(([40],), ([20],)))
2992+
worker._get_kv_split_metadata = MagicMock(
2993+
return_value=(
2994+
[[31001], [31003]],
2995+
[([10],), ([11],)],
2996+
[([30],), ([31],)],
2997+
)
2998+
)
2999+
worker._get_group_pulls_metadata = MagicMock(
3000+
return_value=[
3001+
[[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]],
3002+
[[GroupPull(group_id=0, remote_tp_offset=0, num_group_pulls=1)]],
3003+
]
3004+
)
3005+
worker._get_remote_host_info_by_port = MagicMock(return_value=("localhost", "remote_engine"))
3006+
meta = types.SimpleNamespace(
3007+
remote_request_id="remote_req",
3008+
remote_engine_id="remote_engine",
3009+
remote_host="localhost",
3010+
remote_port=31000,
3011+
remote_pcp_size=2,
3012+
remote_dcp_size=2,
3013+
remote_ptp_size=4,
3014+
remote_multi_nodes_meta_mapping={},
3015+
remote_block_size=16,
3016+
local_block_ids=([10],),
3017+
remote_block_ids=([30],),
3018+
num_computed_tokens=0,
3019+
)
3020+
metadata = types.SimpleNamespace(reqs_in_batch=["req"], requests={"req": meta})
3021+
3022+
worker.start_load_kv(cast(MooncakeConnectorMetadata, metadata))
3023+
3024+
add_request_calls = worker.kv_recv_thread.add_request.call_args_list
3025+
self.assertEqual(len(add_request_calls), 2)
3026+
self.assertEqual(add_request_calls[0].kwargs["remote_handshake_port"], 31001)
3027+
self.assertEqual(add_request_calls[0].kwargs["local_block_ids_replicate_k"], ([40],))
3028+
self.assertEqual(add_request_calls[0].kwargs["remote_block_ids_replicate_k"], ([20],))
3029+
self.assertEqual(add_request_calls[1].kwargs["remote_handshake_port"], 31003)
3030+
self.assertIsNone(add_request_calls[1].kwargs["local_block_ids_replicate_k"])
3031+
self.assertIsNone(add_request_calls[1].kwargs["remote_block_ids_replicate_k"])
3032+
3033+
def test_get_kv_split_metadata_dp1_remote_port_send_num_uses_absolute_ports(self):
3034+
self.vllm_config.kv_transfer_config.kv_port = 30000
3035+
self.vllm_config.model_config.is_deepseek_mla = True
3036+
self.vllm_config.kv_transfer_config.get_from_extra_config.side_effect = lambda k, d: {
3037+
"prefill": {"tp_size": 8, "dp_size": 2, "pp_size": 1},
3038+
"decode": {"tp_size": 4, "dp_size": 4, "pp_size": 1},
3039+
}.get(k, d)
3040+
3041+
worker = MooncakeConnectorWorker(self.vllm_config, self.engine_id, MockKVCacheConfig())
3042+
worker.use_mla = True
3043+
worker.pcp_size = 1
3044+
worker.dcp_size = 1
3045+
worker.tp_size = 4
3046+
worker.tp_rank = 0
3047+
worker.pcp_rank = 0
3048+
worker.dcp_rank = 0
3049+
worker.side_channel_port = 40000
3050+
worker.handshake_port = 40000
3051+
worker.local_remote_block_port_mapping = {}
3052+
worker.remote_port_send_num = {}
3053+
worker.block_size = 16
3054+
worker.num_key_value_heads = 1
3055+
worker.use_sparse = False
3056+
worker.block_size_scale = [[1]]
3057+
worker.kv_group2layeridx = {
3058+
0: (
3059+
{
3060+
"kv_cache_spec_type": "FullAttentionSpec",
3061+
"layer_names": ["model.layers.0.self_attn"],
3062+
},
3063+
[0],
3064+
)
3065+
}
3066+
3067+
remote_mapping = {
3068+
str(offset): {
3069+
"host": f"host-{offset}",
3070+
"engine_id": f"engine-{offset}",
3071+
"handshake_port": 30000 + offset,
3072+
}
3073+
for offset in range(8, 16)
3074+
}
3075+
meta = types.SimpleNamespace(
3076+
remote_pcp_size=1,
3077+
remote_dcp_size=8,
3078+
remote_ptp_size=8,
3079+
remote_port=30008,
3080+
remote_block_ids=(list(range(100, 103)),),
3081+
local_block_ids=(list(range(200, 224)),),
3082+
num_external_tokens=24 * worker.block_size,
3083+
num_prompt_blocks=24,
3084+
num_computed_tokens=0,
3085+
remote_engine_id="remote_engine",
3086+
remote_host="localhost",
3087+
remote_multi_nodes_meta_mapping=remote_mapping,
3088+
remote_block_size=16,
3089+
)
3090+
3091+
ports, _, _ = worker._get_kv_split_metadata("req_dp1", cast(ReqMeta, meta))
3092+
remote_port_send_num = worker.remote_port_send_num[meta.remote_engine_id]
3093+
3094+
self.assertEqual([port for shard in ports for port in shard], list(range(30008, 30016)))
3095+
self.assertEqual(set(remote_port_send_num), set(range(30008, 30016)))
3096+
self.assertNotIn(30016, remote_port_send_num)
3097+
self.assertEqual(remote_port_send_num[30008]["host"], "host-8")
3098+
self.assertEqual(remote_port_send_num[30015]["host"], "host-15")
3099+
self.assertEqual(
3100+
worker._get_remote_host_info_by_port(30008, 30015, "localhost", "remote_engine", remote_mapping),
3101+
("host-15", "engine-15"),
3102+
)
3103+
3104+
def test_get_sfa_replicated_indexer_block_ids_uses_full_blocks_for_prefix(self):
3105+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
3106+
worker.enable_sfa_dcp_replicated_indexer = True
3107+
worker.pcp_size = 1
3108+
worker.dcp_size = 2
3109+
worker.block_size = 16
3110+
meta = types.SimpleNamespace(
3111+
remote_pcp_size=1,
3112+
remote_dcp_size=2,
3113+
remote_block_ids=([10, 11],),
3114+
local_block_ids=([20],),
3115+
local_full_block_ids=([19, 20],),
3116+
num_external_tokens=32,
3117+
num_prompt_blocks=3,
3118+
num_computed_tokens=16,
3119+
)
3120+
3121+
local_ids, remote_ids = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
3122+
3123+
self.assertEqual(local_ids, ([39, 40],))
3124+
self.assertEqual(remote_ids, ([21, 22],))
3125+
3126+
def test_get_sfa_replicated_indexer_block_ids_ignores_empty_regular_kv_shard(self):
3127+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
3128+
worker.enable_sfa_dcp_replicated_indexer = True
3129+
worker.pcp_size = 1
3130+
worker.dcp_size = 2
3131+
worker.block_size = 16
3132+
meta = types.SimpleNamespace(
3133+
remote_pcp_size=1,
3134+
remote_dcp_size=2,
3135+
remote_block_ids=([10],),
3136+
local_block_ids=([],),
3137+
local_full_block_ids=([20],),
3138+
num_external_tokens=16,
3139+
num_prompt_blocks=1,
3140+
num_computed_tokens=0,
3141+
)
3142+
3143+
local_ids, remote_ids = worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
3144+
3145+
self.assertEqual(local_ids, ([40],))
3146+
self.assertEqual(remote_ids, ([20],))
3147+
3148+
def test_get_sfa_replicated_indexer_block_ids_requires_full_blocks_for_prefix(self):
3149+
worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker)
3150+
worker.enable_sfa_dcp_replicated_indexer = True
3151+
worker.pcp_size = 1
3152+
worker.dcp_size = 2
3153+
worker.block_size = 16
3154+
meta = types.SimpleNamespace(
3155+
remote_pcp_size=1,
3156+
remote_dcp_size=2,
3157+
remote_block_ids=([10, 11],),
3158+
local_block_ids=([20],),
3159+
local_full_block_ids=tuple(),
3160+
num_external_tokens=32,
3161+
num_prompt_blocks=3,
3162+
num_computed_tokens=16,
3163+
)
3164+
3165+
with self.assertRaises(AssertionError):
3166+
worker._get_sfa_replicate_k_block_ids(cast(ReqMeta, meta))
3167+
29173168

29183169
if __name__ == "__main__":
29193170
unittest.main()

0 commit comments

Comments
 (0)