@@ -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
14731503class 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
29183169if __name__ == "__main__" :
29193170 unittest .main ()
0 commit comments