[PD] Share the prefill->decode failure notification across backends (#36612)
Co-authored-by: inkcherry <mingzhi.liu@amd.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
inkcherry
Shangming Cai
parent
e016de462c
commit
2a46cf2ca0
@@ -144,6 +144,19 @@ class PrefillRankInfo:
|
|||||||
|
|
||||||
|
|
||||||
class CommonKVManager(BaseKVManager):
|
class CommonKVManager(BaseKVManager):
|
||||||
|
# Wire layout of the prefill->decode terminal status message. The legacy
|
||||||
|
# layout (mooncake, and ascend which inherits it) is three untagged frames
|
||||||
|
# ``[room, status, prefill_rank]``; backends whose control socket also
|
||||||
|
# carries tagged messages prefix a tag frame and may append a reason:
|
||||||
|
# ``[tag, room, status, prefill_rank, reason]``.
|
||||||
|
kv_status_msg_tag: Optional[bytes] = None
|
||||||
|
kv_status_msg_carries_reason: bool = False
|
||||||
|
|
||||||
|
# Used by decode when the prefill reported Failed without a reason frame.
|
||||||
|
DEFAULT_PREFILL_FAILURE_REASON = (
|
||||||
|
"Failed to get kvcache from prefill instance, it might be dead"
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
args: KVArgs,
|
args: KVArgs,
|
||||||
@@ -223,7 +236,6 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self._socket_lock = threading.Lock()
|
self._socket_lock = threading.Lock()
|
||||||
self.failure_records: Dict[int, str] = {}
|
self.failure_records: Dict[int, str] = {}
|
||||||
self.failure_lock = threading.Lock()
|
self.failure_lock = threading.Lock()
|
||||||
|
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
# When SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER is True, all CP ranks
|
# When SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER is True, all CP ranks
|
||||||
# participate in KV transfer; Otherwise only CP rank 0 sends.
|
# participate in KV transfer; Otherwise only CP rank 0 sends.
|
||||||
@@ -253,6 +265,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.bootstrap_timeout = envs.SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT.get()
|
self.bootstrap_timeout = envs.SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT.get()
|
||||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
self.enable_staging: bool = False
|
self.enable_staging: bool = False
|
||||||
|
self._staging_handler = None
|
||||||
self.connection_pool: Dict[str, Dict[str, Union[str, int]]] = {}
|
self.connection_pool: Dict[str, Dict[str, Union[str, int]]] = {}
|
||||||
self.connection_lock = threading.Lock()
|
self.connection_lock = threading.Lock()
|
||||||
self.required_prefill_response_num_table: Dict[int, int] = {}
|
self.required_prefill_response_num_table: Dict[int, int] = {}
|
||||||
@@ -363,26 +376,233 @@ class CommonKVManager(BaseKVManager):
|
|||||||
return self.request_status[bootstrap_room]
|
return self.request_status[bootstrap_room]
|
||||||
|
|
||||||
def update_status(self, bootstrap_room: int, status: KVPoll):
|
def update_status(self, bootstrap_room: int, status: KVPoll):
|
||||||
if bootstrap_room not in self.request_status:
|
current = self.request_status.get(bootstrap_room)
|
||||||
# Do not resurrect a cleared entry with Failed: once clear() has
|
if current is None:
|
||||||
# popped the room from request_status, any late update_status(Failed)
|
# The room does not exist yet, or clear() already popped it. Only a
|
||||||
# (e.g. from abort()) must be a no-op. Otherwise a Failed entry could
|
# request's opening status may create it: Bootstrapping normally, or
|
||||||
# pollute a future request that reuses the same bootstrap_room.
|
# WaitingForInput for a dummy CP rank (see CommonKVSender.__init__).
|
||||||
if status == KVPoll.Failed:
|
# Anything else would resurrect a concluded room and pollute a later
|
||||||
return
|
# request that reuses the same bootstrap_room.
|
||||||
|
if status in (KVPoll.Bootstrapping, KVPoll.WaitingForInput):
|
||||||
self.request_status[bootstrap_room] = status
|
self.request_status[bootstrap_room] = status
|
||||||
else:
|
return
|
||||||
if status == KVPoll.Failed:
|
if status == KVPoll.Failed:
|
||||||
self.request_status[bootstrap_room] = KVPoll.Failed
|
self.request_status[bootstrap_room] = KVPoll.Failed
|
||||||
else:
|
return
|
||||||
self.request_status[bootstrap_room] = max(
|
if current == KVPoll.Failed:
|
||||||
self.request_status[bootstrap_room], status
|
# Failed is terminal. It also sorts lowest, so the max() below would
|
||||||
)
|
# happily promote it back to Transferring or Success.
|
||||||
|
return
|
||||||
|
self.request_status[bootstrap_room] = max(current, status)
|
||||||
|
|
||||||
def record_failure(self, bootstrap_room: int, failure_reason: str):
|
def record_failure(self, bootstrap_room: int, failure_reason: str):
|
||||||
with self.failure_lock:
|
with self.failure_lock:
|
||||||
self.failure_records[bootstrap_room] = failure_reason
|
self.failure_records[bootstrap_room] = failure_reason
|
||||||
|
|
||||||
|
def _room_notify_targets(self, bootstrap_room: int) -> List[Tuple[str, int]]:
|
||||||
|
infos = self.transfer_infos.get(bootstrap_room)
|
||||||
|
if not infos:
|
||||||
|
return []
|
||||||
|
# Every non-dummy endpoint, not just the one a caller failed on: the
|
||||||
|
# others never receive the room's remaining chunks either.
|
||||||
|
targets: List[Tuple[str, int]] = []
|
||||||
|
# Snapshot: the control thread can register a late peer for this room
|
||||||
|
# while we walk it, and iterating the live view would then raise.
|
||||||
|
for info in list(infos.values()):
|
||||||
|
if info.is_dummy:
|
||||||
|
continue
|
||||||
|
target = (info.endpoint, info.dst_port)
|
||||||
|
if target not in targets:
|
||||||
|
targets.append(target)
|
||||||
|
return targets
|
||||||
|
|
||||||
|
def _encode_kv_status_message(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
bootstrap_room: int,
|
||||||
|
status: KVPoll,
|
||||||
|
failure_reason: Optional[str],
|
||||||
|
) -> List[bytes]:
|
||||||
|
parts = [
|
||||||
|
str(bootstrap_room).encode("ascii"),
|
||||||
|
str(int(status)).encode("ascii"),
|
||||||
|
str(self._prefill_unique_rank()).encode("ascii"),
|
||||||
|
]
|
||||||
|
if self.kv_status_msg_carries_reason:
|
||||||
|
parts.append((failure_reason or "").encode("utf-8"))
|
||||||
|
if self.kv_status_msg_tag is not None:
|
||||||
|
parts.insert(0, self.kv_status_msg_tag)
|
||||||
|
return parts
|
||||||
|
|
||||||
|
def parse_kv_status_message(
|
||||||
|
self, msg: List[bytes]
|
||||||
|
) -> Optional[Tuple[int, int, int, Optional[str]]]:
|
||||||
|
"""Decode a prefill status message, or None when it is not one."""
|
||||||
|
if self.kv_status_msg_tag is not None:
|
||||||
|
if not msg or msg[0] != self.kv_status_msg_tag:
|
||||||
|
return None
|
||||||
|
msg = msg[1:]
|
||||||
|
if len(msg) < 3:
|
||||||
|
logger.warning(
|
||||||
|
"Dropping malformed prefill status message with %d frames", len(msg)
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
bootstrap_room = int(msg[0].decode("ascii"))
|
||||||
|
status = int(msg[1].decode("ascii"))
|
||||||
|
prefill_rank = int(msg[2].decode("ascii"))
|
||||||
|
except (UnicodeDecodeError, ValueError):
|
||||||
|
logger.warning("Dropping unparsable prefill status message")
|
||||||
|
return None
|
||||||
|
failure_reason = (
|
||||||
|
msg[3].decode("utf-8", errors="replace")
|
||||||
|
if len(msg) > 3 and msg[3]
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
return bootstrap_room, status, prefill_rank, failure_reason
|
||||||
|
|
||||||
|
def send_kv_status_message(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
targets: List[Tuple[str, int]],
|
||||||
|
bootstrap_room: int,
|
||||||
|
status: KVPoll,
|
||||||
|
failure_reason: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Push of a terminal transfer status to decode endpoints."""
|
||||||
|
if not targets:
|
||||||
|
return
|
||||||
|
parts = self._encode_kv_status_message(
|
||||||
|
bootstrap_room=bootstrap_room,
|
||||||
|
status=status,
|
||||||
|
failure_reason=failure_reason,
|
||||||
|
)
|
||||||
|
for endpoint, dst_port in targets:
|
||||||
|
na = NetworkAddress(endpoint, dst_port)
|
||||||
|
try:
|
||||||
|
self._send_multipart_locked(na.to_tcp(), parts, is_ipv6=na.is_ipv6)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
f"Failed to sync status {status} of room {bootstrap_room} to "
|
||||||
|
f"{na.to_host_port_str()}: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def conclude_transfer(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
bootstrap_room: int,
|
||||||
|
status: KVPoll,
|
||||||
|
targets: Optional[List[Tuple[str, int]]] = None,
|
||||||
|
failure_reason: Optional[str] = None,
|
||||||
|
) -> Optional[KVPoll]:
|
||||||
|
"""Returns the status emitted, or None for a cleared room.
|
||||||
|
|
||||||
|
Runs more than once for a room when a staging chunk is deferred past the
|
||||||
|
last one. ``targets`` defaults to the room's non-dummy decode endpoints.
|
||||||
|
"""
|
||||||
|
if bootstrap_room not in self.request_status:
|
||||||
|
# The sender already cleared this room. Concluding now would
|
||||||
|
# re-create it in request_status and leave a failure record that a
|
||||||
|
# request reusing this bootstrap_room would adopt as its own.
|
||||||
|
return None
|
||||||
|
if status == KVPoll.Success:
|
||||||
|
with self.failure_lock:
|
||||||
|
recorded = self.failure_records.get(bootstrap_room)
|
||||||
|
if recorded is not None:
|
||||||
|
status = KVPoll.Failed
|
||||||
|
failure_reason = recorded
|
||||||
|
elif self.request_status.get(bootstrap_room) == KVPoll.Failed:
|
||||||
|
status = KVPoll.Failed
|
||||||
|
failure_reason = (
|
||||||
|
failure_reason or "Room marked Failed before the transfer ended"
|
||||||
|
)
|
||||||
|
if status == KVPoll.Failed:
|
||||||
|
with self.failure_lock:
|
||||||
|
# Keep the first root cause; later callers see the symptom.
|
||||||
|
failure_reason = self.failure_records.setdefault(
|
||||||
|
bootstrap_room, failure_reason or "KV transfer failed"
|
||||||
|
)
|
||||||
|
|
||||||
|
if targets is None:
|
||||||
|
targets = self._room_notify_targets(bootstrap_room)
|
||||||
|
self.update_status(bootstrap_room, status)
|
||||||
|
self.send_kv_status_message(
|
||||||
|
targets=targets,
|
||||||
|
bootstrap_room=bootstrap_room,
|
||||||
|
status=status,
|
||||||
|
failure_reason=failure_reason,
|
||||||
|
)
|
||||||
|
return status
|
||||||
|
|
||||||
|
def conclude_failure(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
bootstrap_room: int,
|
||||||
|
failure_reason: str,
|
||||||
|
targets: Optional[List[Tuple[str, int]]] = None,
|
||||||
|
) -> Optional[KVPoll]:
|
||||||
|
"""Record the reason, mark the room Failed and tell decode."""
|
||||||
|
return self.conclude_transfer(
|
||||||
|
bootstrap_room=bootstrap_room,
|
||||||
|
status=KVPoll.Failed,
|
||||||
|
targets=targets,
|
||||||
|
failure_reason=failure_reason,
|
||||||
|
)
|
||||||
|
|
||||||
|
def apply_prefill_status(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
bootstrap_room: int,
|
||||||
|
status: int,
|
||||||
|
prefill_rank: int,
|
||||||
|
failure_reason: Optional[str] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Decode-side handling of one prefill rank's terminal status."""
|
||||||
|
if bootstrap_room not in self.request_status:
|
||||||
|
# The room concluded and was cleared. Recording a failure now would
|
||||||
|
# leave an entry that a later request reusing this bootstrap_room
|
||||||
|
# would pick up as its own root cause.
|
||||||
|
logger.debug("Dropping late status for cleared room %s", bootstrap_room)
|
||||||
|
return
|
||||||
|
if status == KVPoll.Success:
|
||||||
|
self.prefill_response_tracker[bootstrap_room].add(prefill_rank)
|
||||||
|
expected_response_num = self.required_prefill_response_num_table.get(
|
||||||
|
bootstrap_room
|
||||||
|
)
|
||||||
|
if expected_response_num is None:
|
||||||
|
logger.warning(
|
||||||
|
"No expected prefill response count for room %s, prefill rank %s",
|
||||||
|
bootstrap_room,
|
||||||
|
prefill_rank,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if (
|
||||||
|
len(self.prefill_response_tracker[bootstrap_room])
|
||||||
|
< expected_response_num
|
||||||
|
):
|
||||||
|
return
|
||||||
|
# Tell the staging handler no more chunks are coming, before any
|
||||||
|
# poller can see Success. Only mooncake gets here: NIXL arms the
|
||||||
|
# handler from its own notifications, mori has no staging.
|
||||||
|
if self.enable_staging and self._staging_handler is not None:
|
||||||
|
handler = self._staging_handler
|
||||||
|
if handler.is_staging_room(bootstrap_room):
|
||||||
|
handler.submit_last_scatter_async(bootstrap_room)
|
||||||
|
self.update_status(bootstrap_room, KVPoll.Success)
|
||||||
|
return
|
||||||
|
if status == KVPoll.Failed:
|
||||||
|
self.record_failure(
|
||||||
|
bootstrap_room, failure_reason or self.DEFAULT_PREFILL_FAILURE_REASON
|
||||||
|
)
|
||||||
|
self.update_status(bootstrap_room, KVPoll.Failed)
|
||||||
|
return
|
||||||
|
logger.warning(
|
||||||
|
"Ignoring non-terminal status %s for room %s from prefill rank %s",
|
||||||
|
status,
|
||||||
|
bootstrap_room,
|
||||||
|
prefill_rank,
|
||||||
|
)
|
||||||
|
|
||||||
def register_deferred_abort_room(self, bootstrap_room: int) -> None:
|
def register_deferred_abort_room(self, bootstrap_room: int) -> None:
|
||||||
"""Arm drain-ack accounting for a held room; a fresh set wipes stale acks
|
"""Arm drain-ack accounting for a held room; a fresh set wipes stale acks
|
||||||
from a prior request that reused this bootstrap_room."""
|
from a prior request that reused this bootstrap_room."""
|
||||||
@@ -1613,6 +1833,9 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.kv_mgr.request_status.pop(self.bootstrap_room, None)
|
self.kv_mgr.request_status.pop(self.bootstrap_room, None)
|
||||||
self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room, None)
|
self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room, None)
|
||||||
self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room, None)
|
self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room, None)
|
||||||
|
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard(
|
||||||
|
self.bootstrap_room
|
||||||
|
)
|
||||||
|
|
||||||
def abort(self):
|
def abort(self):
|
||||||
self.kv_mgr.record_failure(
|
self.kv_mgr.record_failure(
|
||||||
|
|||||||
@@ -951,13 +951,7 @@ def prefetch_staging_reqs(
|
|||||||
full_chunk_pages = staging_grid_tokens(chunked_prefill_size, page_size) // page_size
|
full_chunk_pages = staging_grid_tokens(chunked_prefill_size, page_size) // page_size
|
||||||
|
|
||||||
for session_id, tinfo in transfer_infos[room].items():
|
for session_id, tinfo in transfer_infos[room].items():
|
||||||
# mooncake exposes is_dummy as a dataclass bool field, NIXL exposes it
|
if tinfo.is_dummy:
|
||||||
# as a method (it consults decode_prefix_len). Normalize via callable()
|
|
||||||
# so this shared helper works for either backend; treating a bound
|
|
||||||
# method as truthy (the previous behavior) silently dropped every
|
|
||||||
# STAGING_REQ on NIXL and deadlocked the prefill transfer worker.
|
|
||||||
is_dummy_attr = tinfo.is_dummy
|
|
||||||
if is_dummy_attr() if callable(is_dummy_attr) else is_dummy_attr:
|
|
||||||
continue
|
continue
|
||||||
total_pages = len(tinfo.dst_kv_indices)
|
total_pages = len(tinfo.dst_kv_indices)
|
||||||
if total_pages == 0:
|
if total_pages == 0:
|
||||||
|
|||||||
@@ -292,7 +292,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
self._staging_ctx = DecodeStagingContext() if self.enable_staging else None
|
self._staging_ctx = DecodeStagingContext() if self.enable_staging else None
|
||||||
if self.enable_staging:
|
if self.enable_staging:
|
||||||
self._init_staging_allocator()
|
self._init_staging_allocator()
|
||||||
self._staging_handler = None
|
|
||||||
self.start_decode_thread()
|
self.start_decode_thread()
|
||||||
|
|
||||||
def init_engine(self):
|
def init_engine(self):
|
||||||
@@ -1742,20 +1741,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
|
|
||||||
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
||||||
|
|
||||||
def sync_status_to_decode_endpoint(
|
|
||||||
self, remote: str, dst_port: int, room: int, status: int, prefill_rank: int
|
|
||||||
):
|
|
||||||
na = NetworkAddress(remote, dst_port)
|
|
||||||
self._send_multipart_locked(
|
|
||||||
na.to_tcp(),
|
|
||||||
[
|
|
||||||
str(room).encode("ascii"),
|
|
||||||
str(status).encode("ascii"),
|
|
||||||
str(prefill_rank).encode("ascii"),
|
|
||||||
],
|
|
||||||
is_ipv6=na.is_ipv6,
|
|
||||||
)
|
|
||||||
|
|
||||||
def transfer_worker(
|
def transfer_worker(
|
||||||
self,
|
self,
|
||||||
queue: FastQueue,
|
queue: FastQueue,
|
||||||
@@ -1822,11 +1807,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
polls = []
|
polls = []
|
||||||
dst_ranks_infos = []
|
dst_ranks_infos = []
|
||||||
# Unique id per prefill sender so decode's response set size matches expected_response_num.
|
# Unique id per prefill sender so decode's response set size matches expected_response_num.
|
||||||
prefill_unique_rank = (
|
prefill_unique_rank = self._prefill_unique_rank()
|
||||||
self.attn_tp_rank * (self.pp_size * self.attn_cp_size)
|
|
||||||
+ self.pp_rank * self.attn_cp_size
|
|
||||||
+ self.attn_cp_rank
|
|
||||||
)
|
|
||||||
# When staging transfer is not yet ready (watermark/allocation pending),
|
# When staging transfer is not yet ready (watermark/allocation pending),
|
||||||
# the chunk is re-enqueued and we break out of the req loop to retry later.
|
# the chunk is re-enqueued and we break out of the req loop to retry later.
|
||||||
staging_deferred = False
|
staging_deferred = False
|
||||||
@@ -1836,17 +1817,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
# Early exit if the request has failed
|
# Early exit if the request has failed
|
||||||
with self.session_lock:
|
with self.session_lock:
|
||||||
if req.mooncake_session_id in self.failed_sessions:
|
if req.mooncake_session_id in self.failed_sessions:
|
||||||
self.record_failure(
|
self.conclude_failure(
|
||||||
kv_chunk.room,
|
bootstrap_room=kv_chunk.room,
|
||||||
f"Decode instance could be dead, remote mooncake session {req.mooncake_session_id} is not alive",
|
failure_reason=(
|
||||||
)
|
"Decode instance could be dead, remote "
|
||||||
self.update_status(kv_chunk.room, KVPoll.Failed)
|
f"mooncake session {req.mooncake_session_id} "
|
||||||
self.sync_status_to_decode_endpoint(
|
"is not alive"
|
||||||
req.endpoint,
|
),
|
||||||
req.dst_port,
|
|
||||||
req.room,
|
|
||||||
KVPoll.Failed,
|
|
||||||
prefill_unique_rank,
|
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
@@ -1989,18 +1966,12 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
logger.error(
|
logger.error(
|
||||||
f"Session {req.mooncake_session_id} failed."
|
f"Session {req.mooncake_session_id} failed."
|
||||||
)
|
)
|
||||||
self.record_failure(
|
self.conclude_failure(
|
||||||
kv_chunk.room,
|
bootstrap_room=kv_chunk.room,
|
||||||
|
failure_reason=(
|
||||||
f"Failed to send kv chunk of {kv_chunk.room} to "
|
f"Failed to send kv chunk of {kv_chunk.room} to "
|
||||||
f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}",
|
f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}"
|
||||||
)
|
),
|
||||||
self.update_status(kv_chunk.room, KVPoll.Failed)
|
|
||||||
self.sync_status_to_decode_endpoint(
|
|
||||||
req.endpoint,
|
|
||||||
req.dst_port,
|
|
||||||
req.room,
|
|
||||||
KVPoll.Failed,
|
|
||||||
prefill_unique_rank,
|
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
@@ -2020,18 +1991,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
self.failed_sessions.add(
|
self.failed_sessions.add(
|
||||||
req.mooncake_session_id
|
req.mooncake_session_id
|
||||||
)
|
)
|
||||||
self.record_failure(
|
self.conclude_failure(
|
||||||
kv_chunk.room,
|
bootstrap_room=kv_chunk.room,
|
||||||
f"Failed to send state components of {kv_chunk.room} to "
|
failure_reason=(
|
||||||
f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}",
|
"Failed to send state components of "
|
||||||
)
|
f"{kv_chunk.room} to "
|
||||||
self.update_status(kv_chunk.room, KVPoll.Failed)
|
f"{NetworkAddress(req.endpoint, req.dst_port).to_host_port_str()}"
|
||||||
self.sync_status_to_decode_endpoint(
|
),
|
||||||
req.endpoint,
|
|
||||||
req.dst_port,
|
|
||||||
req.room,
|
|
||||||
KVPoll.Failed,
|
|
||||||
prefill_unique_rank,
|
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
@@ -2042,21 +2008,20 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
target_rank_registration_info.dst_aux_ptrs,
|
target_rank_registration_info.dst_aux_ptrs,
|
||||||
)
|
)
|
||||||
polls.append(True if ret == 0 else False)
|
polls.append(True if ret == 0 else False)
|
||||||
dst_ranks_infos.append(
|
dst_ranks_infos.append((req.endpoint, req.dst_port))
|
||||||
(req.endpoint, req.dst_port, req.room)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Only sync status when all the dst ranks have received the kvcache
|
# Only sync status when all the dst ranks have received the kvcache
|
||||||
if len(polls) == req.required_dst_info_num:
|
if len(polls) == req.required_dst_info_num:
|
||||||
status = KVPoll.Success if all(polls) else KVPoll.Failed
|
status = KVPoll.Success if all(polls) else KVPoll.Failed
|
||||||
self.update_status(req.room, status)
|
self.conclude_transfer(
|
||||||
for endpoint, dst_port, room in dst_ranks_infos:
|
bootstrap_room=req.room,
|
||||||
self.sync_status_to_decode_endpoint(
|
status=status,
|
||||||
endpoint,
|
targets=dst_ranks_infos,
|
||||||
dst_port,
|
failure_reason=(
|
||||||
room,
|
None
|
||||||
status,
|
if status == KVPoll.Success
|
||||||
prefill_unique_rank,
|
else f"Failed to send aux data of {req.room}"
|
||||||
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Dummy request means the decode instance is not used, so its status can be marked as success directly
|
# Dummy request means the decode instance is not used, so its status can be marked as success directly
|
||||||
@@ -2295,32 +2260,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
bootstrap_room, status, prefill_rank = msg
|
parsed = self.parse_kv_status_message(msg)
|
||||||
status = int(status.decode("ascii"))
|
if parsed is None:
|
||||||
bootstrap_room = int(bootstrap_room.decode("ascii"))
|
continue
|
||||||
prefill_rank = int(prefill_rank.decode("ascii"))
|
room, status, prefill_rank, reason = parsed
|
||||||
|
self.apply_prefill_status(
|
||||||
if status == KVPoll.Success:
|
bootstrap_room=room,
|
||||||
if bootstrap_room in self.request_status:
|
status=status,
|
||||||
self.prefill_response_tracker[bootstrap_room].add(prefill_rank)
|
prefill_rank=prefill_rank,
|
||||||
expected_response_num = (
|
failure_reason=reason,
|
||||||
self.required_prefill_response_num_table[bootstrap_room]
|
|
||||||
)
|
)
|
||||||
arrived_response_num = len(
|
|
||||||
self.prefill_response_tracker[bootstrap_room]
|
|
||||||
)
|
|
||||||
if arrived_response_num == expected_response_num:
|
|
||||||
if self.enable_staging:
|
|
||||||
handler = self._staging_handler
|
|
||||||
if handler.is_staging_room(bootstrap_room):
|
|
||||||
handler.submit_last_scatter_async(bootstrap_room)
|
|
||||||
self.update_status(bootstrap_room, KVPoll.Success)
|
|
||||||
elif status == KVPoll.Failed:
|
|
||||||
self.record_failure(
|
|
||||||
bootstrap_room,
|
|
||||||
"Failed to get kvcache from prefill instance, it might be dead",
|
|
||||||
)
|
|
||||||
self.update_status(bootstrap_room, status)
|
|
||||||
|
|
||||||
threading.Thread(target=decode_thread).start()
|
threading.Thread(target=decode_thread).start()
|
||||||
self._start_heartbeat_checker_thread()
|
self._start_heartbeat_checker_thread()
|
||||||
@@ -2380,13 +2329,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
def get_session_id(self):
|
def get_session_id(self):
|
||||||
return self.engine.get_session_id()
|
return self.engine.get_session_id()
|
||||||
|
|
||||||
def _on_heartbeat_success(self, bootstrap_addr: str):
|
|
||||||
current_rooms = self.addr_to_rooms_tracker[bootstrap_addr].copy()
|
|
||||||
for bootstrap_room in current_rooms:
|
|
||||||
# Remove KVPoll.Success requests from the tracker
|
|
||||||
if bootstrap_room not in self.request_status:
|
|
||||||
self.addr_to_rooms_tracker[bootstrap_addr].discard(bootstrap_room)
|
|
||||||
|
|
||||||
def _run_one_probe_pass(self) -> None:
|
def _run_one_probe_pass(self) -> None:
|
||||||
with self.session_lock:
|
with self.session_lock:
|
||||||
snapshot = list(self.failed_sessions)
|
snapshot = list(self.failed_sessions)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import struct
|
|||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Dict, List, Optional, Tuple
|
from typing import List, Optional
|
||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -294,6 +294,11 @@ class TransferTarget:
|
|||||||
class MoriKVManager(CommonKVManager):
|
class MoriKVManager(CommonKVManager):
|
||||||
AUX_DATA_HEADER = b"AUX_DATA"
|
AUX_DATA_HEADER = b"AUX_DATA"
|
||||||
|
|
||||||
|
# The bootstrap socket carries several message kinds, so the status message
|
||||||
|
# is tagged. Mori has always shipped the failure reason with it.
|
||||||
|
kv_status_msg_tag = MORI_GUARD
|
||||||
|
kv_status_msg_carries_reason = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
args: KVArgs,
|
args: KVArgs,
|
||||||
@@ -319,8 +324,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
]
|
]
|
||||||
self._wait_poll_ms = envs.SGLANG_MORI_WAIT_POLL_MS.get()
|
self._wait_poll_ms = envs.SGLANG_MORI_WAIT_POLL_MS.get()
|
||||||
self._transfer_timeout_ms = envs.SGLANG_MORI_TRANSFER_TIMEOUT_MS.get()
|
self._transfer_timeout_ms = envs.SGLANG_MORI_TRANSFER_TIMEOUT_MS.get()
|
||||||
self._room_status_notified: Dict[int, bool] = {}
|
|
||||||
self._room_notify_lock = threading.Lock()
|
|
||||||
for shard, queue in enumerate(self._transfer_queues):
|
for shard, queue in enumerate(self._transfer_queues):
|
||||||
threading.Thread(
|
threading.Thread(
|
||||||
target=self._transfer_worker,
|
target=self._transfer_worker,
|
||||||
@@ -333,7 +336,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
).start()
|
).start()
|
||||||
self._start_bootstrap_thread()
|
self._start_bootstrap_thread()
|
||||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
self.room_to_bootstrap_addr: Dict[int, str] = {}
|
|
||||||
self._start_decode_thread()
|
self._start_decode_thread()
|
||||||
self._start_heartbeat_checker_thread()
|
self._start_heartbeat_checker_thread()
|
||||||
|
|
||||||
@@ -415,19 +417,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
component_descs.append(desc)
|
component_descs.append(desc)
|
||||||
self.state_mem_descs.append(component_descs)
|
self.state_mem_descs.append(component_descs)
|
||||||
|
|
||||||
def update_status(self, bootstrap_room: int, status: KVPoll):
|
|
||||||
current = self.request_status.get(bootstrap_room)
|
|
||||||
if current is None:
|
|
||||||
# Room not yet created or already cleared.
|
|
||||||
# Only allow initial creation: Bootstrapping (normal) or
|
|
||||||
# WaitingForInput (dummy CP rank, see CommonKVSender.__init__).
|
|
||||||
if status not in (KVPoll.Bootstrapping, KVPoll.WaitingForInput):
|
|
||||||
return
|
|
||||||
elif current == KVPoll.Failed and status != KVPoll.Failed:
|
|
||||||
# Failed is terminal — never overwrite with non-Failed.
|
|
||||||
return
|
|
||||||
super().update_status(bootstrap_room, status)
|
|
||||||
|
|
||||||
def _transfer_worker(self, queue: FastQueue) -> None:
|
def _transfer_worker(self, queue: FastQueue) -> None:
|
||||||
while True:
|
while True:
|
||||||
kv_chunk = queue.get()
|
kv_chunk = queue.get()
|
||||||
@@ -443,7 +432,9 @@ class MoriKVManager(CommonKVManager):
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
try:
|
try:
|
||||||
self._conclude_room_failure(kv_chunk.room, failure_reason)
|
self.conclude_failure(
|
||||||
|
bootstrap_room=kv_chunk.room, failure_reason=failure_reason
|
||||||
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
try:
|
try:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
@@ -464,7 +455,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
if self._should_skip_transfer(room):
|
if self._should_skip_transfer(room):
|
||||||
return
|
return
|
||||||
|
|
||||||
statuses, target_infos = self._submit_kv_transfer(
|
statuses = self._submit_kv_transfer(
|
||||||
room,
|
room,
|
||||||
kv_chunk.prefill_kv_indices,
|
kv_chunk.prefill_kv_indices,
|
||||||
kv_chunk.index_slice,
|
kv_chunk.index_slice,
|
||||||
@@ -480,14 +471,14 @@ class MoriKVManager(CommonKVManager):
|
|||||||
if self._should_skip_transfer(room):
|
if self._should_skip_transfer(room):
|
||||||
return
|
return
|
||||||
if failure_reason is not None:
|
if failure_reason is not None:
|
||||||
self._conclude_room_failure(room, failure_reason)
|
self.conclude_failure(bootstrap_room=room, failure_reason=failure_reason)
|
||||||
return
|
return
|
||||||
|
|
||||||
if kv_chunk.is_last_chunk:
|
if kv_chunk.is_last_chunk:
|
||||||
self._notify_decode_for_room(
|
# conclude_transfer downgrades to Failed when a failure was recorded
|
||||||
room, KVPoll.Success, target_infos=target_infos
|
# while this chunk was in flight, and applies the same status locally
|
||||||
)
|
# and on the wire.
|
||||||
self.update_status(room, KVPoll.Success)
|
self.conclude_transfer(bootstrap_room=room, status=KVPoll.Success)
|
||||||
|
|
||||||
def _should_skip_transfer(self, room: int) -> bool:
|
def _should_skip_transfer(self, room: int) -> bool:
|
||||||
if room not in self.request_status or self.check_status(room) == KVPoll.Failed:
|
if room not in self.request_status or self.check_status(room) == KVPoll.Failed:
|
||||||
@@ -523,60 +514,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
return f"KV transfer failed: {status.Message()}"
|
return f"KV transfer failed: {status.Message()}"
|
||||||
return "KV transfer failed due to unknown reason"
|
return "KV transfer failed due to unknown reason"
|
||||||
|
|
||||||
def _notify_decode_for_room(
|
|
||||||
self,
|
|
||||||
room: int,
|
|
||||||
status: KVPoll,
|
|
||||||
failure_reason: Optional[str] = None,
|
|
||||||
target_infos: Optional[List[TransferInfo]] = None,
|
|
||||||
) -> None:
|
|
||||||
with self._room_notify_lock:
|
|
||||||
if room not in self.request_status or self._room_status_notified.get(room):
|
|
||||||
return
|
|
||||||
|
|
||||||
emitted_status = status
|
|
||||||
emitted_reason = failure_reason
|
|
||||||
|
|
||||||
if emitted_status == KVPoll.Success:
|
|
||||||
with self.failure_lock:
|
|
||||||
recorded = self.failure_records.get(room)
|
|
||||||
if recorded is not None:
|
|
||||||
emitted_status = KVPoll.Failed
|
|
||||||
emitted_reason = recorded
|
|
||||||
elif self.request_status.get(room) == KVPoll.Failed:
|
|
||||||
emitted_status = KVPoll.Failed
|
|
||||||
emitted_reason = (
|
|
||||||
emitted_reason or "request marked Failed before notify"
|
|
||||||
)
|
|
||||||
|
|
||||||
if emitted_status == KVPoll.Failed:
|
|
||||||
with self.failure_lock:
|
|
||||||
self.failure_records.setdefault(
|
|
||||||
room, emitted_reason or "KV transfer failed"
|
|
||||||
)
|
|
||||||
self.update_status(room, KVPoll.Failed)
|
|
||||||
|
|
||||||
infos = target_infos
|
|
||||||
if infos is None:
|
|
||||||
with self.transfer_lock:
|
|
||||||
room_infos = self.transfer_infos.get(room)
|
|
||||||
infos = (
|
|
||||||
list(room_infos.values()) if room_infos is not None else None
|
|
||||||
)
|
|
||||||
|
|
||||||
self._room_status_notified[room] = True
|
|
||||||
|
|
||||||
if infos:
|
|
||||||
self.notify_decode_status(infos, room, emitted_status, emitted_reason)
|
|
||||||
|
|
||||||
def _conclude_room_failure(
|
|
||||||
self, room: int, failure_reason: Optional[str] = None
|
|
||||||
) -> None:
|
|
||||||
if failure_reason is None:
|
|
||||||
with self.failure_lock:
|
|
||||||
failure_reason = self.failure_records.get(room, "KV transfer failed")
|
|
||||||
self._notify_decode_for_room(room, KVPoll.Failed, failure_reason)
|
|
||||||
|
|
||||||
def add_transfer_request(
|
def add_transfer_request(
|
||||||
self,
|
self,
|
||||||
bootstrap_room: int,
|
bootstrap_room: int,
|
||||||
@@ -772,15 +709,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
|
|
||||||
threading.Thread(target=bootstrap_worker, daemon=True).start()
|
threading.Thread(target=bootstrap_worker, daemon=True).start()
|
||||||
|
|
||||||
def _cleanup_room_tracking(self, bootstrap_room: int) -> None:
|
|
||||||
bootstrap_addr = self.room_to_bootstrap_addr.pop(bootstrap_room, None)
|
|
||||||
if bootstrap_addr is not None:
|
|
||||||
rooms = self.addr_to_rooms_tracker.get(bootstrap_addr)
|
|
||||||
if rooms is not None:
|
|
||||||
rooms.discard(bootstrap_room)
|
|
||||||
if not rooms:
|
|
||||||
self.addr_to_rooms_tracker.pop(bootstrap_addr, None)
|
|
||||||
|
|
||||||
def _start_decode_thread(self) -> None:
|
def _start_decode_thread(self) -> None:
|
||||||
def decode_worker():
|
def decode_worker():
|
||||||
while True:
|
while True:
|
||||||
@@ -790,97 +718,24 @@ class MoriKVManager(CommonKVManager):
|
|||||||
self._handle_aux_data(msg)
|
self._handle_aux_data(msg)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not msg or msg[0] != MORI_GUARD:
|
parsed = self.parse_kv_status_message(msg)
|
||||||
|
if parsed is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Received malformed status message on decode worker"
|
"Received malformed status message on decode worker"
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
payload = msg[1:]
|
room, status, prefill_rank, reason = parsed
|
||||||
if len(payload) < 3:
|
self.apply_prefill_status(
|
||||||
logger.warning("Incomplete status payload received")
|
bootstrap_room=room,
|
||||||
continue
|
status=status,
|
||||||
bootstrap_room = int(payload[0].decode("ascii"))
|
prefill_rank=prefill_rank,
|
||||||
if bootstrap_room not in self.request_status:
|
failure_reason=reason,
|
||||||
logger.debug(
|
|
||||||
"Dropping late status for cleared room %s",
|
|
||||||
bootstrap_room,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
status_code = int(payload[1].decode("ascii"))
|
|
||||||
prefill_rank = int(payload[2].decode("ascii"))
|
|
||||||
failure_reason = (
|
|
||||||
payload[3].decode("utf-8")
|
|
||||||
if len(payload) > 3 and payload[3]
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
if status_code == KVPoll.Success:
|
|
||||||
tracker = self.prefill_response_tracker[bootstrap_room]
|
|
||||||
tracker.add(prefill_rank)
|
|
||||||
expected = self.required_prefill_response_num_table.get(
|
|
||||||
bootstrap_room, 1
|
|
||||||
)
|
|
||||||
if len(tracker) >= expected:
|
|
||||||
self.prefill_response_tracker.pop(bootstrap_room, None)
|
|
||||||
self.update_status(bootstrap_room, KVPoll.Success)
|
|
||||||
self._cleanup_room_tracking(bootstrap_room)
|
|
||||||
elif status_code == KVPoll.Failed:
|
|
||||||
if failure_reason:
|
|
||||||
self.record_failure(bootstrap_room, failure_reason)
|
|
||||||
self.prefill_response_tracker.pop(bootstrap_room, None)
|
|
||||||
self.update_status(bootstrap_room, KVPoll.Failed)
|
|
||||||
self._cleanup_room_tracking(bootstrap_room)
|
|
||||||
else:
|
|
||||||
logger.warning(
|
|
||||||
"Unknown status code %s received for room %s",
|
|
||||||
status_code,
|
|
||||||
bootstrap_room,
|
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Decode status worker failed")
|
logger.exception("Decode status worker failed")
|
||||||
|
|
||||||
threading.Thread(target=decode_worker, daemon=True).start()
|
threading.Thread(target=decode_worker, daemon=True).start()
|
||||||
|
|
||||||
def _compute_prefill_unique_rank(self) -> int:
|
|
||||||
"""Unique id per prefill sender, encoding TP/PP/CP ranks.
|
|
||||||
Must match Mooncake's formula so decode's response set size matches
|
|
||||||
expected_response_num when multiple CP ranks participate."""
|
|
||||||
return (
|
|
||||||
self.attn_tp_rank * (self.pp_size * self.attn_cp_size)
|
|
||||||
+ self.pp_rank * self.attn_cp_size
|
|
||||||
+ self.attn_cp_rank
|
|
||||||
)
|
|
||||||
|
|
||||||
def notify_decode_status(
|
|
||||||
self,
|
|
||||||
infos: List[TransferInfo],
|
|
||||||
bootstrap_room: int,
|
|
||||||
status: KVPoll,
|
|
||||||
failure_reason: Optional[str] = None,
|
|
||||||
) -> None:
|
|
||||||
if not infos:
|
|
||||||
return
|
|
||||||
payload = [
|
|
||||||
MORI_GUARD,
|
|
||||||
str(bootstrap_room).encode("ascii"),
|
|
||||||
str(int(status)).encode("ascii"),
|
|
||||||
str(self._compute_prefill_unique_rank()).encode("ascii"),
|
|
||||||
failure_reason.encode("utf-8") if failure_reason else b"",
|
|
||||||
]
|
|
||||||
for info in infos:
|
|
||||||
try:
|
|
||||||
na = NetworkAddress(info.endpoint, info.dst_port)
|
|
||||||
socket = self._connect_threadsafe(na.to_tcp(), is_ipv6=na.is_ipv6)
|
|
||||||
socket.send_multipart(payload)
|
|
||||||
except Exception:
|
|
||||||
logger.exception(
|
|
||||||
"Failed to sync status %s to decode endpoint %s:%s for room %s",
|
|
||||||
status,
|
|
||||||
info.endpoint,
|
|
||||||
info.dst_port,
|
|
||||||
bootstrap_room,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _add_remote_peer(self, register_info: KVArgsRegisterInfo) -> None:
|
def _add_remote_peer(self, register_info: KVArgsRegisterInfo) -> None:
|
||||||
engine_key = register_info.engine_key
|
engine_key = register_info.engine_key
|
||||||
if engine_key in self.decode_kv_args_table:
|
if engine_key in self.decode_kv_args_table:
|
||||||
@@ -1541,21 +1396,20 @@ class MoriKVManager(CommonKVManager):
|
|||||||
is_last_chunk: bool,
|
is_last_chunk: bool,
|
||||||
aux_index: Optional[int] = None,
|
aux_index: Optional[int] = None,
|
||||||
state_indices: Optional[List[npt.NDArray[np.int32]]] = None,
|
state_indices: Optional[List[npt.NDArray[np.int32]]] = None,
|
||||||
) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]:
|
) -> List[TransferStatus]:
|
||||||
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||||
|
|
||||||
if (
|
if (
|
||||||
bootstrap_room not in self.request_status
|
bootstrap_room not in self.request_status
|
||||||
or self.request_status.get(bootstrap_room) == KVPoll.Failed
|
or self.request_status.get(bootstrap_room) == KVPoll.Failed
|
||||||
):
|
):
|
||||||
return [], None
|
return []
|
||||||
|
|
||||||
targets: List[TransferTarget] = []
|
targets: List[TransferTarget] = []
|
||||||
target_infos_snapshot: Optional[List[TransferInfo]] = None
|
|
||||||
with self.transfer_lock:
|
with self.transfer_lock:
|
||||||
current = self.request_status.get(bootstrap_room)
|
current = self.request_status.get(bootstrap_room)
|
||||||
if current is None or current == KVPoll.Failed:
|
if current is None or current == KVPoll.Failed:
|
||||||
return [], None
|
return []
|
||||||
|
|
||||||
transfer_infos = self.transfer_infos.get(bootstrap_room)
|
transfer_infos = self.transfer_infos.get(bootstrap_room)
|
||||||
if not transfer_infos:
|
if not transfer_infos:
|
||||||
@@ -1571,8 +1425,6 @@ class MoriKVManager(CommonKVManager):
|
|||||||
f"Peer info missing for engine {info.engine_key}"
|
f"Peer info missing for engine {info.engine_key}"
|
||||||
)
|
)
|
||||||
targets.append(TransferTarget(info=info, peer_info=peer_info))
|
targets.append(TransferTarget(info=info, peer_info=peer_info))
|
||||||
if is_last_chunk:
|
|
||||||
target_infos_snapshot = list(transfer_infos.values())
|
|
||||||
|
|
||||||
result_statuses: List[TransferStatus] = []
|
result_statuses: List[TransferStatus] = []
|
||||||
try:
|
try:
|
||||||
@@ -1616,7 +1468,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
raise RuntimeError(f"Transfer submission failed: {e}") from e
|
raise RuntimeError(f"Transfer submission failed: {e}") from e
|
||||||
|
|
||||||
return result_statuses, target_infos_snapshot
|
return result_statuses
|
||||||
|
|
||||||
|
|
||||||
class MoriKVSender(CommonKVSender):
|
class MoriKVSender(CommonKVSender):
|
||||||
@@ -1705,11 +1557,6 @@ class MoriKVSender(CommonKVSender):
|
|||||||
self.conclude_state = status
|
self.conclude_state = status
|
||||||
return status
|
return status
|
||||||
|
|
||||||
def clear(self) -> None:
|
|
||||||
super().clear()
|
|
||||||
with self.kv_mgr._room_notify_lock:
|
|
||||||
self.kv_mgr._room_status_notified.pop(self.bootstrap_room, None)
|
|
||||||
|
|
||||||
def failure_exception(self):
|
def failure_exception(self):
|
||||||
if self.conclude_state is None:
|
if self.conclude_state is None:
|
||||||
self.conclude_state = KVPoll.Failed
|
self.conclude_state = KVPoll.Failed
|
||||||
@@ -1741,9 +1588,6 @@ class MoriKVReceiver(CommonKVReceiver):
|
|||||||
prefill_dp_rank: int,
|
prefill_dp_rank: int,
|
||||||
):
|
):
|
||||||
super().init(prefill_dp_rank)
|
super().init(prefill_dp_rank)
|
||||||
if self.bootstrap_room is None:
|
|
||||||
return
|
|
||||||
self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr
|
|
||||||
|
|
||||||
def _register_kv_args(self) -> bool:
|
def _register_kv_args(self) -> bool:
|
||||||
if self.bootstrap_infos is None:
|
if self.bootstrap_infos is None:
|
||||||
@@ -1871,7 +1715,6 @@ class MoriKVReceiver(CommonKVReceiver):
|
|||||||
if self.bootstrap_room is None:
|
if self.bootstrap_room is None:
|
||||||
return
|
return
|
||||||
super().clear()
|
super().clear()
|
||||||
self.kv_mgr._cleanup_room_tracking(self.bootstrap_room)
|
|
||||||
|
|
||||||
def failure_exception(self):
|
def failure_exception(self):
|
||||||
if self.conclude_state is None:
|
if self.conclude_state is None:
|
||||||
|
|||||||
@@ -72,6 +72,12 @@ logger = logging.getLogger(__name__)
|
|||||||
GUARD = "NixlMsgGuard".encode("ascii")
|
GUARD = "NixlMsgGuard".encode("ascii")
|
||||||
KV_MEM_KINDS = {"VRAM", "DRAM"}
|
KV_MEM_KINDS = {"VRAM", "DRAM"}
|
||||||
|
|
||||||
|
# Once one handle of a batch reports ERR, its siblings settle only when NIXL
|
||||||
|
# notices their peer is gone, which for UCX means waiting out peer keepalive.
|
||||||
|
# This worker serves other rooms, so bound that wait.
|
||||||
|
NIXL_ERR_SETTLE_TIMEOUT_S = 5.0
|
||||||
|
NIXL_ERR_SETTLE_POLL_S = 0.001
|
||||||
|
|
||||||
|
|
||||||
def _normalize_kv_mem_kinds(kinds: Optional[List[str]], expected_len: int) -> List[str]:
|
def _normalize_kv_mem_kinds(kinds: Optional[List[str]], expected_len: int) -> List[str]:
|
||||||
if kinds is None:
|
if kinds is None:
|
||||||
@@ -166,45 +172,47 @@ class TransferInfo:
|
|||||||
required_dst_info_num: int
|
required_dst_info_num: int
|
||||||
dst_state_indices: List[List[int]]
|
dst_state_indices: List[List[int]]
|
||||||
decode_prefix_len: Optional[int] = None # for decode radix cache
|
decode_prefix_len: Optional[int] = None # for decode radix cache
|
||||||
is_dummy_rank: Optional[bool] = None
|
is_dummy: bool = False
|
||||||
# NOTE: optional staging field; populated via STAGING_RSP. Keep at the
|
# NOTE: optional staging field; populated via STAGING_RSP. Keep at the
|
||||||
# end so positional construction in from_zmq() continues to work.
|
# end so positional construction in from_zmq() continues to work.
|
||||||
staging: Optional[StagingTransferInfo] = None
|
staging: Optional[StagingTransferInfo] = None
|
||||||
|
|
||||||
def is_dummy(self):
|
|
||||||
# A transfer is "dummy" only for CP non-authoritative ranks.
|
|
||||||
# When dst_kv_indices is empty due to a decode-side radix cache
|
|
||||||
# full hit (decode_prefix_len > 0), the transfer is NOT dummy --
|
|
||||||
# aux/state data still needs to be sent.
|
|
||||||
if self.is_dummy_rank is not None:
|
|
||||||
return self.is_dummy_rank
|
|
||||||
if self.dst_kv_indices.size == 0 and self.decode_prefix_len:
|
|
||||||
return False
|
|
||||||
return self.dst_kv_indices.size == 0
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
dst_state_indices = (
|
dst_state_indices = (
|
||||||
unpack_int_lists(msg[7], "i") if len(msg) > 7 and msg[7] != b"" else []
|
unpack_int_lists(msg[7], "i") if len(msg) > 7 and msg[7] != b"" else []
|
||||||
)
|
)
|
||||||
|
dst_kv_indices = np.frombuffer(msg[4], dtype=np.int32)
|
||||||
|
decode_prefix_len = (
|
||||||
|
int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None
|
||||||
|
) # hacky just add it into the message that will be sent
|
||||||
|
dummy_rank = (
|
||||||
|
bool(int(msg[9].decode("ascii")))
|
||||||
|
if len(msg) > 9 and msg[9] != b""
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
# A transfer is "dummy" only for CP non-authoritative ranks. When
|
||||||
|
# dst_kv_indices is empty due to a decode-side radix cache full hit
|
||||||
|
# (decode_prefix_len > 0), the transfer is NOT dummy -- aux/state data
|
||||||
|
# still needs to be sent.
|
||||||
|
if dummy_rank is not None:
|
||||||
|
is_dummy = dummy_rank
|
||||||
|
elif dst_kv_indices.size == 0 and decode_prefix_len:
|
||||||
|
is_dummy = False
|
||||||
|
else:
|
||||||
|
is_dummy = dst_kv_indices.size == 0
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
room=int(msg[0].decode("ascii")),
|
room=int(msg[0].decode("ascii")),
|
||||||
endpoint=msg[1].decode("ascii"),
|
endpoint=msg[1].decode("ascii"),
|
||||||
dst_port=int(msg[2].decode("ascii")),
|
dst_port=int(msg[2].decode("ascii")),
|
||||||
agent_name=msg[3].decode("ascii"),
|
agent_name=msg[3].decode("ascii"),
|
||||||
dst_kv_indices=np.frombuffer(msg[4], dtype=np.int32),
|
dst_kv_indices=dst_kv_indices,
|
||||||
dst_aux_index=int(msg[5].decode("ascii")),
|
dst_aux_index=int(msg[5].decode("ascii")),
|
||||||
required_dst_info_num=int(msg[6].decode("ascii")),
|
required_dst_info_num=int(msg[6].decode("ascii")),
|
||||||
dst_state_indices=dst_state_indices,
|
dst_state_indices=dst_state_indices,
|
||||||
decode_prefix_len=(
|
decode_prefix_len=decode_prefix_len,
|
||||||
int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None
|
is_dummy=is_dummy,
|
||||||
), # hacky just add it into the message that will be sent
|
|
||||||
is_dummy_rank=(
|
|
||||||
bool(int(msg[9].decode("ascii")))
|
|
||||||
if len(msg) > 9 and msg[9] != b""
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -399,6 +407,11 @@ class TransferStatus:
|
|||||||
|
|
||||||
|
|
||||||
class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||||
|
# The decode control socket multiplexes tagged messages, so the status
|
||||||
|
# message is tagged too. It is new to NIXL, hence free to carry the reason.
|
||||||
|
kv_status_msg_tag = b"KV_STATUS"
|
||||||
|
kv_status_msg_carries_reason = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
args: KVArgs,
|
args: KVArgs,
|
||||||
@@ -523,8 +536,6 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
)
|
)
|
||||||
if self.enable_staging:
|
if self.enable_staging:
|
||||||
self._init_staging_decode_ctx()
|
self._init_staging_decode_ctx()
|
||||||
self._staging_handler = None
|
|
||||||
if self.enable_staging or self.enable_deferred_decode_kv_release:
|
|
||||||
self._start_decode_listener_thread()
|
self._start_decode_listener_thread()
|
||||||
self._start_heartbeat_checker_thread()
|
self._start_heartbeat_checker_thread()
|
||||||
else:
|
else:
|
||||||
@@ -595,14 +606,18 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
self._staging_ctx.room_receivers[room] = receiver
|
self._staging_ctx.room_receivers[room] = receiver
|
||||||
|
|
||||||
def _start_decode_listener_thread(self):
|
def _start_decode_listener_thread(self):
|
||||||
"""Decode-side ZMQ listener for STAGING_REQ and ABORT_ACK. A thread, not
|
"""Decode-side ZMQ listener for KV_STATUS, STAGING_REQ and ABORT_ACK. A
|
||||||
NIXL notifs: the decode agent has no progress thread, so notifs only drain
|
thread, not NIXL notifs: the decode agent has no progress thread, so notifs
|
||||||
inside a live receiver's poll() and would be missed while idle."""
|
only drain inside a live receiver's poll() and would be missed while idle.
|
||||||
|
|
||||||
|
Started unconditionally: KV_STATUS carries prefill-side transfer failures,
|
||||||
|
which are independent of staging and deferred KV release."""
|
||||||
|
|
||||||
def decode_listener_thread():
|
def decode_listener_thread():
|
||||||
while True:
|
while True:
|
||||||
msg = self.server_socket.recv_multipart()
|
msg = self.server_socket.recv_multipart()
|
||||||
if msg[0] == b"STAGING_REQ":
|
if msg[0] == b"STAGING_REQ":
|
||||||
|
if self.enable_staging:
|
||||||
self._handle_staging_req(msg)
|
self._handle_staging_req(msg)
|
||||||
continue
|
continue
|
||||||
if msg[0] == b"ABORT_ACK":
|
if msg[0] == b"ABORT_ACK":
|
||||||
@@ -612,6 +627,16 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
int(msg[1].decode("ascii")), int(msg[2].decode("ascii"))
|
int(msg[1].decode("ascii")), int(msg[2].decode("ascii"))
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
parsed = self.parse_kv_status_message(msg)
|
||||||
|
if parsed is not None:
|
||||||
|
room, status, prefill_rank, reason = parsed
|
||||||
|
self.apply_prefill_status(
|
||||||
|
bootstrap_room=room,
|
||||||
|
status=status,
|
||||||
|
prefill_rank=prefill_rank,
|
||||||
|
failure_reason=reason,
|
||||||
|
)
|
||||||
|
continue
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"decode_listener_thread: unexpected message tag %s",
|
"decode_listener_thread: unexpected message tag %s",
|
||||||
msg[0][:20],
|
msg[0][:20],
|
||||||
@@ -634,7 +659,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
|
|
||||||
room_infos = self.transfer_infos.get(room, {})
|
room_infos = self.transfer_infos.get(room, {})
|
||||||
needs_staging = any(
|
needs_staging = any(
|
||||||
not tinfo.is_dummy()
|
not tinfo.is_dummy
|
||||||
and tinfo.agent_name in self.decode_kv_args_table
|
and tinfo.agent_name in self.decode_kv_args_table
|
||||||
and self.decode_kv_args_table[tinfo.agent_name].decode_tp_size
|
and self.decode_kv_args_table[tinfo.agent_name].decode_tp_size
|
||||||
!= self.attn_tp_size
|
!= self.attn_tp_size
|
||||||
@@ -662,11 +687,44 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
def check_status(self, bootstrap_room: int):
|
def check_status(self, bootstrap_room: int):
|
||||||
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
|
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
|
||||||
|
|
||||||
def update_status(self, bootstrap_room: int, status: KVPoll):
|
def _await_handles(
|
||||||
# Keep Failed sticky until the sender clears the room.
|
self, handles: List[Any], *, failure_seen: bool
|
||||||
if self.request_status.get(bootstrap_room) == KVPoll.Failed:
|
) -> Tuple[bool, bool]:
|
||||||
return
|
"""Poll until every handle settled. Returns ``(settled, any_failed)``.
|
||||||
super().update_status(bootstrap_room, status)
|
|
||||||
|
The wait is unbounded while every handle is still healthy, and bounded
|
||||||
|
to NIXL_ERR_SETTLE_TIMEOUT_S from the moment the batch is known broken.
|
||||||
|
``failure_seen`` arms that deadline up front, for a batch that raised
|
||||||
|
before the barrier ran. A state that cannot be read counts as running,
|
||||||
|
since it does not prove the write into the decode's pages is over.
|
||||||
|
"""
|
||||||
|
deadline = time.time() + NIXL_ERR_SETTLE_TIMEOUT_S if failure_seen else None
|
||||||
|
while True:
|
||||||
|
all_settled = True
|
||||||
|
any_failed = failure_seen
|
||||||
|
try:
|
||||||
|
for handle in handles:
|
||||||
|
state = self.agent.check_xfer_state(handle)
|
||||||
|
if state == "ERR":
|
||||||
|
any_failed = True
|
||||||
|
elif state != "DONE":
|
||||||
|
all_settled = False
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Failed to read NIXL transfer state: {e}")
|
||||||
|
return False, True
|
||||||
|
if all_settled:
|
||||||
|
return True, any_failed
|
||||||
|
if not any_failed:
|
||||||
|
time.sleep(0)
|
||||||
|
continue
|
||||||
|
# This room is already lost, so trade its notification for the
|
||||||
|
# worker's other rooms: back off, and give up waiting for the
|
||||||
|
# siblings once the deadline passes.
|
||||||
|
if deadline is None:
|
||||||
|
deadline = time.time() + NIXL_ERR_SETTLE_TIMEOUT_S
|
||||||
|
elif time.time() >= deadline:
|
||||||
|
return False, True
|
||||||
|
time.sleep(NIXL_ERR_SETTLE_POLL_S)
|
||||||
|
|
||||||
def _prep_equal_tp_dlist(
|
def _prep_equal_tp_dlist(
|
||||||
self,
|
self,
|
||||||
@@ -1095,6 +1153,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
kv_chunk: TransferKVChunk = queue.get()
|
kv_chunk: TransferKVChunk = queue.get()
|
||||||
room = kv_chunk.room
|
room = kv_chunk.room
|
||||||
handles: List[Any] = []
|
handles: List[Any] = []
|
||||||
|
settle_timed_out = False
|
||||||
try:
|
try:
|
||||||
if room not in self.request_status:
|
if room not in self.request_status:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
@@ -1154,7 +1213,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
|
|
||||||
for req in reqs_to_be_processed:
|
for req in reqs_to_be_processed:
|
||||||
assert room == req.room
|
assert room == req.room
|
||||||
if req.is_dummy():
|
if req.is_dummy:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
assert req.agent_name in self.decode_kv_args_table
|
assert req.agent_name in self.decode_kv_args_table
|
||||||
@@ -1353,19 +1412,19 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
# Chunk has been re-enqueued; do not advance status.
|
# Chunk has been re-enqueued; do not advance status.
|
||||||
continue
|
continue
|
||||||
|
|
||||||
while handles:
|
# Raise only once every handle of this batch settled, not on the
|
||||||
all_done = True
|
# first ERR: a sibling still in PROC keeps writing into the
|
||||||
for handle in handles:
|
# decode's KV pages, and the failure path below tells the decode
|
||||||
state = self.agent.check_xfer_state(handle)
|
# those pages are free.
|
||||||
if state == "ERR":
|
settled, any_failed = self._await_handles(handles, failure_seen=False)
|
||||||
|
if not settled:
|
||||||
|
settle_timed_out = True
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"NIXL transfer encountered ERR room={room}"
|
f"NIXL transfer for room {room} left a handle running "
|
||||||
|
f"{NIXL_ERR_SETTLE_TIMEOUT_S}s after a peer handle failed"
|
||||||
)
|
)
|
||||||
if state != "DONE":
|
if any_failed:
|
||||||
all_done = False
|
raise RuntimeError(f"NIXL transfer encountered ERR room={room}")
|
||||||
if all_done:
|
|
||||||
break
|
|
||||||
time.sleep(0)
|
|
||||||
|
|
||||||
self._staging_outstanding[room] -= 1
|
self._staging_outstanding[room] -= 1
|
||||||
if self.enable_deferred_decode_kv_release:
|
if self.enable_deferred_decode_kv_release:
|
||||||
@@ -1409,10 +1468,20 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
f"Unexpected transfer worker error for room {room}"
|
f"Unexpected transfer worker error for room {room}"
|
||||||
)
|
)
|
||||||
self.exceptions[room] = e
|
self.exceptions[room] = e
|
||||||
|
# An exception raised while the batch was still being built
|
||||||
|
# leaves the handles posted so far running, so settle here too
|
||||||
|
# rather than only after the barrier.
|
||||||
|
notify = False
|
||||||
|
if not settle_timed_out:
|
||||||
|
notify, _ = self._await_handles(handles, failure_seen=True)
|
||||||
|
if notify:
|
||||||
|
self.conclude_failure(bootstrap_room=room, failure_reason=str(e))
|
||||||
|
else:
|
||||||
|
# A handle can still write into the decode's KV pages, so
|
||||||
|
# leave the room to the decode's waiting timeout rather
|
||||||
|
# than telling it those pages are free.
|
||||||
self.record_failure(room, str(e))
|
self.record_failure(room, str(e))
|
||||||
self.update_status(room, KVPoll.Failed)
|
self.update_status(room, KVPoll.Failed)
|
||||||
# No ack here on purpose: the DONE barrier bails on the first
|
|
||||||
# ERR, so siblings may still be writing; fall back to the timeout.
|
|
||||||
|
|
||||||
def register_buffer_to_engine(self):
|
def register_buffer_to_engine(self):
|
||||||
self.kv_descs = []
|
self.kv_descs = []
|
||||||
@@ -3002,6 +3071,13 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
super().__init__(mgr, bootstrap_addr, bootstrap_room)
|
super().__init__(mgr, bootstrap_addr, bootstrap_room)
|
||||||
self.init_time = None
|
self.init_time = None
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
super().clear()
|
||||||
|
# transfer_statuses is NIXL's own per-room bookkeeping -- the other
|
||||||
|
# backends track completion through prefill_response_tracker, which
|
||||||
|
# CommonKVReceiver.clear() already drops -- so it needs its own pop.
|
||||||
|
self.kv_mgr.transfer_statuses.pop(self.bootstrap_room, None)
|
||||||
|
|
||||||
def send_metadata(
|
def send_metadata(
|
||||||
self,
|
self,
|
||||||
kv_indices: npt.NDArray[np.int32],
|
kv_indices: npt.NDArray[np.int32],
|
||||||
@@ -3096,11 +3172,7 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
# deadline would otherwise lose to the timeout purely by poll ordering.
|
# deadline would otherwise lose to the timeout purely by poll ordering.
|
||||||
self.kv_mgr.update_transfer_status()
|
self.kv_mgr.update_transfer_status()
|
||||||
if self.kv_mgr.check_transfer_done(self.bootstrap_room): # type: ignore
|
if self.kv_mgr.check_transfer_done(self.bootstrap_room): # type: ignore
|
||||||
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].discard(
|
|
||||||
self.bootstrap_room
|
|
||||||
)
|
|
||||||
self.conclude_state = KVPoll.Success
|
self.conclude_state = KVPoll.Success
|
||||||
del self.kv_mgr.transfer_statuses[self.bootstrap_room]
|
|
||||||
return self.conclude_state # type: ignore
|
return self.conclude_state # type: ignore
|
||||||
|
|
||||||
timeout_result = self._check_waiting_timeout()
|
timeout_result = self._check_waiting_timeout()
|
||||||
@@ -3201,6 +3273,11 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def failure_exception(self):
|
def failure_exception(self):
|
||||||
|
if self.conclude_state is None:
|
||||||
|
self.conclude_state = KVPoll.Failed
|
||||||
|
|
||||||
|
self.clear()
|
||||||
|
|
||||||
with self.kv_mgr.failure_lock:
|
with self.kv_mgr.failure_lock:
|
||||||
failure_reason = self.kv_mgr.failure_records.pop(self.bootstrap_room, None)
|
failure_reason = self.kv_mgr.failure_records.pop(self.bootstrap_room, None)
|
||||||
is_propagated = failure_reason is None
|
is_propagated = failure_reason is None
|
||||||
|
|||||||
@@ -165,7 +165,7 @@ class TestNixlTransferInfo(CustomTestCase):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertFalse(info.is_dummy())
|
self.assertFalse(info.is_dummy)
|
||||||
|
|
||||||
def test_empty_indices_without_decode_prefix_is_dummy(self):
|
def test_empty_indices_without_decode_prefix_is_dummy(self):
|
||||||
info = TransferInfo.from_zmq(
|
info = TransferInfo.from_zmq(
|
||||||
@@ -182,7 +182,7 @@ class TestNixlTransferInfo(CustomTestCase):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertTrue(info.is_dummy())
|
self.assertTrue(info.is_dummy)
|
||||||
|
|
||||||
def test_explicit_dummy_frame_true_is_dummy(self):
|
def test_explicit_dummy_frame_true_is_dummy(self):
|
||||||
# msg[9] is the explicit is_dummy frame the sender writes
|
# msg[9] is the explicit is_dummy frame the sender writes
|
||||||
@@ -848,10 +848,9 @@ class TestNixlReceiverPoll(CustomTestCase):
|
|||||||
mgr.update_transfer_status.assert_called_once_with()
|
mgr.update_transfer_status.assert_called_once_with()
|
||||||
mgr.record_failure.assert_not_called()
|
mgr.record_failure.assert_not_called()
|
||||||
mgr.update_status.assert_not_called()
|
mgr.update_status.assert_not_called()
|
||||||
self.assertNotIn(11, mgr.transfer_statuses)
|
|
||||||
|
|
||||||
@patch("sglang.srt.disaggregation.nixl.conn.time.time")
|
@patch("sglang.srt.disaggregation.nixl.conn.time.time")
|
||||||
def test_transfer_done_returns_success_and_cleans_room_state(self, mock_time):
|
def test_transfer_done_returns_success_and_clear_drops_room_state(self, mock_time):
|
||||||
mock_time.return_value = 12.0
|
mock_time.return_value = 12.0
|
||||||
receiver, mgr = self._make_receiver(status=KVPoll.WaitingForInput)
|
receiver, mgr = self._make_receiver(status=KVPoll.WaitingForInput)
|
||||||
receiver.started_transfer = True
|
receiver.started_transfer = True
|
||||||
@@ -864,9 +863,16 @@ class TestNixlReceiverPoll(CustomTestCase):
|
|||||||
mgr.check_transfer_done.return_value = True
|
mgr.check_transfer_done.return_value = True
|
||||||
|
|
||||||
self.assertEqual(receiver.poll(), KVPoll.Success)
|
self.assertEqual(receiver.poll(), KVPoll.Success)
|
||||||
|
self.assertEqual(receiver.conclude_state, KVPoll.Success)
|
||||||
|
|
||||||
|
# poll() only concludes now; dropping room state is left to clear(), the
|
||||||
|
# way mooncake and mori already do it. The scheduler calls clear() as
|
||||||
|
# soon as poll() reports Success or Failed, so both terminal paths clean
|
||||||
|
# up -- the cleanup that used to live in poll() ran on Success only.
|
||||||
|
receiver.clear()
|
||||||
|
|
||||||
self.assertNotIn(11, mgr.transfer_statuses)
|
self.assertNotIn(11, mgr.transfer_statuses)
|
||||||
self.assertNotIn(11, mgr.addr_to_rooms_tracker["prefill:8998"])
|
self.assertNotIn(11, mgr.addr_to_rooms_tracker["prefill:8998"])
|
||||||
self.assertEqual(receiver.conclude_state, KVPoll.Success)
|
|
||||||
|
|
||||||
|
|
||||||
class TestNixlNodeFailure(CustomTestCase):
|
class TestNixlNodeFailure(CustomTestCase):
|
||||||
|
|||||||
Reference in New Issue
Block a user