[PD] Support decode pp for PD disaggregation (#14265)

Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Shangming Cai
2025-12-03 14:35:29 +08:00
committed by GitHub
parent 65c8568c4a
commit 93452a8252
2 changed files with 29 additions and 13 deletions
@@ -153,13 +153,17 @@ class CommonKVManager(BaseKVManager):
def get_mha_kv_ptrs_with_pp( def get_mha_kv_ptrs_with_pp(
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int] self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
) -> Tuple[List[int], List[int], List[int], List[int], int]: ) -> Tuple[List[int], List[int], List[int], List[int], int]:
# pp is not supported on the decode side yet
start_layer = self.kv_args.prefill_start_layer start_layer = self.kv_args.prefill_start_layer
num_kv_layers = len(src_kv_ptrs) // 2 num_kv_layers = len(src_kv_ptrs) // 2
end_layer = start_layer + num_kv_layers end_layer = start_layer + num_kv_layers
dst_num_total_layers = len(dst_kv_ptrs) // 2 dst_num_total_layers = len(dst_kv_ptrs) // 2
src_k_ptrs = src_kv_ptrs[:num_kv_layers] src_k_ptrs = src_kv_ptrs[:num_kv_layers]
src_v_ptrs = src_kv_ptrs[num_kv_layers:] src_v_ptrs = src_kv_ptrs[num_kv_layers:]
if num_kv_layers == dst_num_total_layers:
dst_k_ptrs = dst_kv_ptrs[:dst_num_total_layers]
dst_v_ptrs = dst_kv_ptrs[dst_num_total_layers:]
else:
# Decode pp size should be equal to prefill pp size or 1
dst_k_ptrs = dst_kv_ptrs[start_layer:end_layer] dst_k_ptrs = dst_kv_ptrs[start_layer:end_layer]
dst_v_ptrs = dst_kv_ptrs[ dst_v_ptrs = dst_kv_ptrs[
dst_num_total_layers + start_layer : dst_num_total_layers + end_layer dst_num_total_layers + start_layer : dst_num_total_layers + end_layer
@@ -170,9 +174,12 @@ class CommonKVManager(BaseKVManager):
def get_mla_kv_ptrs_with_pp( def get_mla_kv_ptrs_with_pp(
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int] self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
) -> Tuple[List[int], List[int], int]: ) -> Tuple[List[int], List[int], int]:
# pp is not supported on the decode side yet
start_layer = self.kv_args.prefill_start_layer start_layer = self.kv_args.prefill_start_layer
end_layer = start_layer + len(src_kv_ptrs) end_layer = start_layer + len(src_kv_ptrs)
if len(src_kv_ptrs) == len(dst_kv_ptrs):
sliced_dst_kv_ptrs = dst_kv_ptrs
else:
# Decode pp size should be equal to prefill pp size or 1
sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer] sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer]
layers_current_pp_stage = len(src_kv_ptrs) layers_current_pp_stage = len(src_kv_ptrs)
return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage
@@ -273,8 +280,7 @@ class CommonKVReceiver(BaseKVReceiver):
self.bootstrap_addr self.bootstrap_addr
] ]
# Currently, we don't allow prefill instance and decode instance to # Handling for PD with different TP sizes per DP rank
# have different TP sizes per DP rank, except for models using MLA.
if self.kv_mgr.attn_tp_size == self.prefill_attn_tp_size: if self.kv_mgr.attn_tp_size == self.prefill_attn_tp_size:
self.target_tp_rank = ( self.target_tp_rank = (
self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size
@@ -335,9 +341,19 @@ class CommonKVReceiver(BaseKVReceiver):
else: else:
self.prefill_dp_rank = bootstrap_room % self.prefill_dp_size self.prefill_dp_rank = bootstrap_room % self.prefill_dp_size
# FIXME: alias here: target_dp_group -> prefill_dp_rank
self.target_dp_group = self.prefill_dp_rank self.target_dp_group = self.prefill_dp_rank
# Decode pp size should be equal to prefill pp size or 1
assert (
self.kv_mgr.pp_size == self.prefill_pp_size or self.kv_mgr.pp_size == 1
), (
f"Decode pp size ({self.kv_mgr.pp_size}) should be equal to prefill pp size ({self.prefill_pp_size}) or 1",
)
if self.prefill_pp_size == self.kv_mgr.pp_size:
self.target_pp_ranks = [self.kv_mgr.pp_rank]
else:
self.target_pp_ranks = [rank for rank in range(self.prefill_pp_size)]
self.kv_mgr.required_prefill_response_num_table[self.bootstrap_room] = ( self.kv_mgr.required_prefill_response_num_table[self.bootstrap_room] = (
self.required_prefill_response_num self.required_prefill_response_num
) )
@@ -349,7 +365,7 @@ class CommonKVReceiver(BaseKVReceiver):
if bootstrap_key not in self.kv_mgr.connection_pool: if bootstrap_key not in self.kv_mgr.connection_pool:
bootstrap_infos = [] bootstrap_infos = []
for target_tp_rank in self.target_tp_ranks: for target_tp_rank in self.target_tp_ranks:
for target_pp_rank in range(self.prefill_pp_size): for target_pp_rank in self.target_pp_ranks:
bootstrap_info = self._get_bootstrap_info_from_server( bootstrap_info = self._get_bootstrap_info_from_server(
target_tp_rank, self.target_dp_group, target_pp_rank target_tp_rank, self.target_dp_group, target_pp_rank
) )
@@ -283,7 +283,7 @@ class MooncakeKVManager(CommonKVManager):
layers_params = None layers_params = None
# pp is not supported on the decode side yet # Decode pp size should be equal to prefill pp size or 1
if self.is_mla_backend: if self.is_mla_backend:
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = ( src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs) self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs)
@@ -1199,7 +1199,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
packed_state_data_ptrs = b"".join( packed_state_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs
) )
# Note(shangming): No need to add pp rank here since pp is not supported on the decode side yet # Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1
tp_rank = self.kv_mgr.kv_args.engine_rank tp_rank = self.kv_mgr.kv_args.engine_rank
kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0] kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
dst_tp_rank = str(tp_rank).encode("ascii") dst_tp_rank = str(tp_rank).encode("ascii")