[PD] Add kv_cache_dtype consistency check for PD Disaggregation (#19407)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -51,6 +51,7 @@ class PrefillServerInfo:
|
|||||||
dp_size: int
|
dp_size: int
|
||||||
pp_size: int
|
pp_size: int
|
||||||
page_size: Optional[int]
|
page_size: Optional[int]
|
||||||
|
kv_cache_dtype: Optional[str]
|
||||||
follow_bootstrap_room: bool
|
follow_bootstrap_room: bool
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
@@ -58,6 +59,9 @@ class PrefillServerInfo:
|
|||||||
self.dp_size = int(self.dp_size)
|
self.dp_size = int(self.dp_size)
|
||||||
self.pp_size = int(self.pp_size)
|
self.pp_size = int(self.pp_size)
|
||||||
self.page_size = int(self.page_size) if self.page_size is not None else None
|
self.page_size = int(self.page_size) if self.page_size is not None else None
|
||||||
|
self.kv_cache_dtype = (
|
||||||
|
str(self.kv_cache_dtype) if self.kv_cache_dtype is not None else None
|
||||||
|
)
|
||||||
self.follow_bootstrap_room = bool(self.follow_bootstrap_room)
|
self.follow_bootstrap_room = bool(self.follow_bootstrap_room)
|
||||||
|
|
||||||
|
|
||||||
@@ -187,6 +191,16 @@ class CommonKVManager(BaseKVManager):
|
|||||||
f"Both servers must use the same --page-size value."
|
f"Both servers must use the same --page-size value."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
info.kv_cache_dtype is not None
|
||||||
|
and info.kv_cache_dtype != self.server_args.kv_cache_dtype
|
||||||
|
):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, "
|
||||||
|
f"but decode server has kv_cache_dtype={self.server_args.kv_cache_dtype}. "
|
||||||
|
f"Both servers must use the same --kv-cache-dtype value."
|
||||||
|
)
|
||||||
|
|
||||||
self.prefill_info_table[bootstrap_addr] = info
|
self.prefill_info_table[bootstrap_addr] = info
|
||||||
logger.debug(f"Prefill parallel info for [{bootstrap_addr}]: {info}")
|
logger.debug(f"Prefill parallel info for [{bootstrap_addr}]: {info}")
|
||||||
return True
|
return True
|
||||||
@@ -242,6 +256,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"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,
|
"page_size": self.kv_args.page_size,
|
||||||
|
"kv_cache_dtype": self.server_args.kv_cache_dtype,
|
||||||
"load_balance_method": self.server_args.load_balance_method,
|
"load_balance_method": self.server_args.load_balance_method,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -603,6 +618,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
self.attn_tp_size = None
|
self.attn_tp_size = None
|
||||||
self.dp_size = dp_size
|
self.dp_size = dp_size
|
||||||
self.page_size = None
|
self.page_size = None
|
||||||
|
self.kv_cache_dtype: Optional[str] = None
|
||||||
self.follow_bootstrap_room: Optional[bool] = None
|
self.follow_bootstrap_room: Optional[bool] = 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]]]]
|
||||||
@@ -662,6 +678,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
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"])
|
page_size = int(data["page_size"])
|
||||||
|
kv_cache_dtype = data["kv_cache_dtype"]
|
||||||
|
|
||||||
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
|
||||||
@@ -675,6 +692,9 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
if self.page_size is None and page_size is not None:
|
if self.page_size is None and page_size is not None:
|
||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
|
|
||||||
|
if self.kv_cache_dtype is None and kv_cache_dtype is not None:
|
||||||
|
self.kv_cache_dtype = kv_cache_dtype
|
||||||
|
|
||||||
if self.follow_bootstrap_room is None:
|
if self.follow_bootstrap_room is None:
|
||||||
load_balance_method = data.get(
|
load_balance_method = data.get(
|
||||||
"load_balance_method", "follow_bootstrap_room"
|
"load_balance_method", "follow_bootstrap_room"
|
||||||
@@ -730,6 +750,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
dp_size=self.dp_size,
|
dp_size=self.dp_size,
|
||||||
pp_size=self.pp_size,
|
pp_size=self.pp_size,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
follow_bootstrap_room=(
|
follow_bootstrap_room=(
|
||||||
self.follow_bootstrap_room
|
self.follow_bootstrap_room
|
||||||
if self.follow_bootstrap_room is not None
|
if self.follow_bootstrap_room is not None
|
||||||
|
|||||||
Reference in New Issue
Block a user