[PD] Support decode pp for PD disaggregation (#14265)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user