feat: raise error in PD when page sizes are mismatched (#14474)
Signed-off-by: Raayan Dhar raayan.dhar@gmail.com <raayan.dhar@gmail.com> Signed-off-by: raayandhar <raayan.dhar@gmail.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Shangming Cai
parent
d56fd10cb6
commit
4397cda7dc
@@ -92,6 +92,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.prefill_attn_tp_size_table: Dict[str, int] = {}
|
self.prefill_attn_tp_size_table: Dict[str, int] = {}
|
||||||
self.prefill_dp_size_table: Dict[str, int] = {}
|
self.prefill_dp_size_table: Dict[str, int] = {}
|
||||||
self.prefill_pp_size_table: Dict[str, int] = {}
|
self.prefill_pp_size_table: Dict[str, int] = {}
|
||||||
|
self.prefill_page_size_table: Dict[str, Optional[int]] = {}
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
||||||
@@ -127,6 +128,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"system_dp_rank": self.system_dp_rank,
|
"system_dp_rank": self.system_dp_rank,
|
||||||
"rank_ip": self.local_ip,
|
"rank_ip": self.local_ip,
|
||||||
"rank_port": self.rank_port,
|
"rank_port": self.rank_port,
|
||||||
|
"page_size": self.kv_args.page_size,
|
||||||
}
|
}
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -243,6 +245,7 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.prefill_attn_tp_size,
|
self.prefill_attn_tp_size,
|
||||||
self.prefill_dp_size,
|
self.prefill_dp_size,
|
||||||
self.prefill_pp_size,
|
self.prefill_pp_size,
|
||||||
|
self.prefill_page_size,
|
||||||
) = self._get_prefill_parallel_info_from_server()
|
) = self._get_prefill_parallel_info_from_server()
|
||||||
if (
|
if (
|
||||||
self.prefill_attn_tp_size is None
|
self.prefill_attn_tp_size is None
|
||||||
@@ -256,9 +259,20 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
|
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
|
||||||
self.bootstrap_infos = None
|
self.bootstrap_infos = None
|
||||||
return
|
return
|
||||||
else:
|
|
||||||
|
if self.prefill_page_size is not None:
|
||||||
|
decode_page_size = self.kv_mgr.kv_args.page_size
|
||||||
|
if self.prefill_page_size != decode_page_size:
|
||||||
|
error_msg = (
|
||||||
|
f"Page size mismatch: prefill server has page_size={self.prefill_page_size}, "
|
||||||
|
f"but decode server has page_size={decode_page_size}. "
|
||||||
|
f"Both servers must use the same --page-size value."
|
||||||
|
)
|
||||||
|
logger.error(error_msg)
|
||||||
|
raise RuntimeError(error_msg)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Fetch prefill parallel info from [{self.bootstrap_addr}]: DP size:{self.prefill_dp_size}, TP size:{self.prefill_attn_tp_size} PP size:{self.prefill_pp_size}"
|
f"Fetch prefill parallel info from [{self.bootstrap_addr}]: DP size:{self.prefill_dp_size}, TP size:{self.prefill_attn_tp_size} PP size:{self.prefill_pp_size} Page size:{self.prefill_page_size}"
|
||||||
)
|
)
|
||||||
self.kv_mgr.prefill_attn_tp_size_table[self.bootstrap_addr] = (
|
self.kv_mgr.prefill_attn_tp_size_table[self.bootstrap_addr] = (
|
||||||
self.prefill_attn_tp_size
|
self.prefill_attn_tp_size
|
||||||
@@ -269,6 +283,9 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.kv_mgr.prefill_pp_size_table[self.bootstrap_addr] = (
|
self.kv_mgr.prefill_pp_size_table[self.bootstrap_addr] = (
|
||||||
self.prefill_pp_size
|
self.prefill_pp_size
|
||||||
)
|
)
|
||||||
|
self.kv_mgr.prefill_page_size_table[self.bootstrap_addr] = (
|
||||||
|
self.prefill_page_size
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
self.prefill_attn_tp_size = self.kv_mgr.prefill_attn_tp_size_table[
|
self.prefill_attn_tp_size = self.kv_mgr.prefill_attn_tp_size_table[
|
||||||
self.bootstrap_addr
|
self.bootstrap_addr
|
||||||
@@ -279,6 +296,9 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.prefill_pp_size = self.kv_mgr.prefill_pp_size_table[
|
self.prefill_pp_size = self.kv_mgr.prefill_pp_size_table[
|
||||||
self.bootstrap_addr
|
self.bootstrap_addr
|
||||||
]
|
]
|
||||||
|
self.prefill_page_size = self.kv_mgr.prefill_page_size_table.get(
|
||||||
|
self.bootstrap_addr
|
||||||
|
)
|
||||||
|
|
||||||
# Handling for PD with different TP sizes per DP rank
|
# Handling for PD with different TP sizes per DP rank
|
||||||
if self.kv_mgr.attn_tp_size == self.prefill_attn_tp_size:
|
if self.kv_mgr.attn_tp_size == self.prefill_attn_tp_size:
|
||||||
@@ -423,7 +443,7 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
|
|
||||||
def _get_prefill_parallel_info_from_server(
|
def _get_prefill_parallel_info_from_server(
|
||||||
self,
|
self,
|
||||||
) -> Tuple[Optional[int], Optional[int], Optional[int]]:
|
) -> Tuple[Optional[int], Optional[int], Optional[int], Optional[int]]:
|
||||||
"""Fetch the prefill parallel info from the bootstrap server."""
|
"""Fetch the prefill parallel info from the bootstrap server."""
|
||||||
try:
|
try:
|
||||||
url = f"http://{self.bootstrap_addr}/route?engine_rank={-1}&target_dp_group={-1}&target_pp_rank={-1}"
|
url = f"http://{self.bootstrap_addr}/route?engine_rank={-1}&target_dp_group={-1}&target_pp_rank={-1}"
|
||||||
@@ -434,15 +454,16 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
int(prefill_parallel_info["prefill_attn_tp_size"]),
|
int(prefill_parallel_info["prefill_attn_tp_size"]),
|
||||||
int(prefill_parallel_info["prefill_dp_size"]),
|
int(prefill_parallel_info["prefill_dp_size"]),
|
||||||
int(prefill_parallel_info["prefill_pp_size"]),
|
int(prefill_parallel_info["prefill_pp_size"]),
|
||||||
|
int(prefill_parallel_info["prefill_page_size"]),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Failed to get prefill parallel info: {response.status_code}, {response.text}"
|
f"Failed to get prefill parallel info: {response.status_code}, {response.text}"
|
||||||
)
|
)
|
||||||
return None, None, None
|
return None, None, None, None
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error fetching prefill parallel info from bootstrap: {e}")
|
logger.error(f"Error fetching prefill parallel info from bootstrap: {e}")
|
||||||
return None, None, None
|
return None, None, None, None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _connect(cls, endpoint: str, is_ipv6: bool = False):
|
def _connect(cls, endpoint: str, is_ipv6: bool = False):
|
||||||
@@ -484,6 +505,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
self.pp_size = None
|
self.pp_size = None
|
||||||
self.attn_tp_size = None
|
self.attn_tp_size = None
|
||||||
self.dp_size = None
|
self.dp_size = None
|
||||||
|
self.page_size = None
|
||||||
self.prefill_port_table: Dict[
|
self.prefill_port_table: Dict[
|
||||||
int, Dict[int, Dict[int, Dict[str, Union[str, int]]]]
|
int, Dict[int, Dict[int, Dict[str, Union[str, int]]]]
|
||||||
] = {}
|
] = {}
|
||||||
@@ -526,6 +548,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
system_dp_rank = data["system_dp_rank"]
|
system_dp_rank = data["system_dp_rank"]
|
||||||
rank_ip = data["rank_ip"]
|
rank_ip = data["rank_ip"]
|
||||||
rank_port = int(data["rank_port"])
|
rank_port = int(data["rank_port"])
|
||||||
|
page_size = int(data["page_size"])
|
||||||
|
|
||||||
if self.attn_tp_size is None:
|
if self.attn_tp_size is None:
|
||||||
self.attn_tp_size = attn_tp_size
|
self.attn_tp_size = attn_tp_size
|
||||||
@@ -536,6 +559,9 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
if self.pp_size is None:
|
if self.pp_size is None:
|
||||||
self.pp_size = pp_size
|
self.pp_size = pp_size
|
||||||
|
|
||||||
|
if self.page_size is None and page_size is not None:
|
||||||
|
self.page_size = page_size
|
||||||
|
|
||||||
if role == "Prefill":
|
if role == "Prefill":
|
||||||
if system_dp_size == 1:
|
if system_dp_size == 1:
|
||||||
dp_group = attn_dp_rank
|
dp_group = attn_dp_rank
|
||||||
@@ -576,6 +602,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
"prefill_attn_tp_size": self.attn_tp_size,
|
"prefill_attn_tp_size": self.attn_tp_size,
|
||||||
"prefill_dp_size": self.dp_size,
|
"prefill_dp_size": self.dp_size,
|
||||||
"prefill_pp_size": self.pp_size,
|
"prefill_pp_size": self.pp_size,
|
||||||
|
"prefill_page_size": self.page_size,
|
||||||
}
|
}
|
||||||
return web.json_response(prefill_parallel_info, status=200)
|
return web.json_response(prefill_parallel_info, status=200)
|
||||||
|
|
||||||
|
|||||||
@@ -274,6 +274,7 @@ class DecodePreallocQueue:
|
|||||||
kv_args.kv_data_ptrs = kv_data_ptrs
|
kv_args.kv_data_ptrs = kv_data_ptrs
|
||||||
kv_args.kv_data_lens = kv_data_lens
|
kv_args.kv_data_lens = kv_data_lens
|
||||||
kv_args.kv_item_lens = kv_item_lens
|
kv_args.kv_item_lens = kv_item_lens
|
||||||
|
kv_args.page_size = self.token_to_kv_pool.page_size
|
||||||
|
|
||||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
||||||
self.metadata_buffers.get_buf_infos()
|
self.metadata_buffers.get_buf_infos()
|
||||||
|
|||||||
Reference in New Issue
Block a user