[PD][mori] Align prefill transfer control plane for unified control plane (#36160)
This commit is contained in:
@@ -32,6 +32,8 @@ class TransferKVChunk:
|
|||||||
# Set when the staging worker first counts this chunk toward the per-room
|
# Set when the staging worker first counts this chunk toward the per-room
|
||||||
# outstanding count; stays set across re-enqueue on a watermark defer.
|
# outstanding count; stays set across re-enqueue on a watermark defer.
|
||||||
staging_counted: bool = False
|
staging_counted: bool = False
|
||||||
|
# Mori early-send: CUDA event to synchronize before RDMA (optional).
|
||||||
|
wait_event: Optional[object] = None
|
||||||
|
|
||||||
|
|
||||||
def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
|
def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from sglang.srt.disaggregation.common.conn import (
|
|||||||
from sglang.srt.disaggregation.common.utils import (
|
from sglang.srt.disaggregation.common.utils import (
|
||||||
AuxDataCodec,
|
AuxDataCodec,
|
||||||
FastQueue,
|
FastQueue,
|
||||||
|
TransferKVChunk,
|
||||||
group_concurrent_contiguous,
|
group_concurrent_contiguous,
|
||||||
pack_int_lists,
|
pack_int_lists,
|
||||||
unpack_int_lists,
|
unpack_int_lists,
|
||||||
@@ -288,17 +289,6 @@ class TransferTarget:
|
|||||||
peer_info: KVArgsRegisterInfo
|
peer_info: KVArgsRegisterInfo
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
|
||||||
class _TransferChunk:
|
|
||||||
sender: MoriKVSender
|
|
||||||
kv_indices: npt.NDArray[np.int32]
|
|
||||||
index_slice: slice
|
|
||||||
is_last_chunk: bool
|
|
||||||
aux_index: Optional[int]
|
|
||||||
normalized_state: Optional[List[Optional[npt.NDArray[np.int32]]]]
|
|
||||||
wait_event: Optional[object] = None
|
|
||||||
|
|
||||||
|
|
||||||
class MoriKVManager(CommonKVManager):
|
class MoriKVManager(CommonKVManager):
|
||||||
AUX_DATA_HEADER = b"AUX_DATA"
|
AUX_DATA_HEADER = b"AUX_DATA"
|
||||||
|
|
||||||
@@ -327,6 +317,8 @@ 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,
|
||||||
@@ -429,34 +421,195 @@ class MoriKVManager(CommonKVManager):
|
|||||||
return
|
return
|
||||||
super().update_status(bootstrap_room, status)
|
super().update_status(bootstrap_room, status)
|
||||||
|
|
||||||
def enqueue_transfer(self, task: _TransferChunk) -> None:
|
|
||||||
self._transfer_queues[task.sender.bootstrap_room % self._num_shards].put(task)
|
|
||||||
|
|
||||||
def _transfer_worker(self, queue: FastQueue) -> None:
|
def _transfer_worker(self, queue: FastQueue) -> None:
|
||||||
while True:
|
while True:
|
||||||
task = queue.get()
|
kv_chunk = queue.get()
|
||||||
try:
|
try:
|
||||||
task.sender._run_chunk(task)
|
self._process_transfer_chunk(kv_chunk)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
failure_reason = f"transfer worker raised: {exc!r}"
|
failure_reason = f"transfer worker raised: {exc!r}"
|
||||||
try:
|
try:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
"Mori transfer worker failed for room %s",
|
"Mori transfer worker failed for room %s",
|
||||||
task.sender.bootstrap_room,
|
kv_chunk.room,
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
try:
|
try:
|
||||||
task.sender._fail_from_worker(failure_reason)
|
self._conclude_room_failure(kv_chunk.room, failure_reason)
|
||||||
except Exception:
|
except Exception:
|
||||||
try:
|
try:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
"Mori transfer worker failover failed for room %s",
|
"Mori transfer worker failover failed for room %s",
|
||||||
task.sender.bootstrap_room,
|
kv_chunk.room,
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def _process_transfer_chunk(self, kv_chunk: TransferKVChunk) -> None:
|
||||||
|
room = kv_chunk.room
|
||||||
|
if self._should_skip_transfer(room):
|
||||||
|
return
|
||||||
|
|
||||||
|
if kv_chunk.wait_event is not None:
|
||||||
|
kv_chunk.wait_event.synchronize()
|
||||||
|
|
||||||
|
if self._should_skip_transfer(room):
|
||||||
|
return
|
||||||
|
|
||||||
|
statuses, target_infos = self._submit_kv_transfer(
|
||||||
|
room,
|
||||||
|
kv_chunk.prefill_kv_indices,
|
||||||
|
kv_chunk.index_slice,
|
||||||
|
kv_chunk.is_last_chunk,
|
||||||
|
aux_index=kv_chunk.prefill_aux_index,
|
||||||
|
state_indices=kv_chunk.state_indices,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self._should_skip_transfer(room):
|
||||||
|
return
|
||||||
|
|
||||||
|
failure_reason = self._wait_transfer_completion(statuses)
|
||||||
|
if self._should_skip_transfer(room):
|
||||||
|
return
|
||||||
|
if failure_reason is not None:
|
||||||
|
self._conclude_room_failure(room, failure_reason)
|
||||||
|
return
|
||||||
|
|
||||||
|
if kv_chunk.is_last_chunk:
|
||||||
|
self._notify_decode_for_room(
|
||||||
|
room, KVPoll.Success, target_infos=target_infos
|
||||||
|
)
|
||||||
|
self.update_status(room, KVPoll.Success)
|
||||||
|
|
||||||
|
def _should_skip_transfer(self, room: int) -> bool:
|
||||||
|
if room not in self.request_status or self.check_status(room) == KVPoll.Failed:
|
||||||
|
logger.debug(
|
||||||
|
"Skipping chunk for room %s because it has already failed or been aborted",
|
||||||
|
room,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _wait_transfer_completion(
|
||||||
|
self, statuses: List[TransferStatus]
|
||||||
|
) -> Optional[str]:
|
||||||
|
if not statuses:
|
||||||
|
return None
|
||||||
|
|
||||||
|
start = time.perf_counter()
|
||||||
|
sla_ms = self._transfer_timeout_ms
|
||||||
|
|
||||||
|
while True:
|
||||||
|
rc = self.engine.wait_all(statuses, timeout_ms=self._wait_poll_ms)
|
||||||
|
if rc != StatusCode.IN_PROGRESS:
|
||||||
|
if rc == StatusCode.SUCCESS:
|
||||||
|
return None
|
||||||
|
return self._collect_transfer_failure_reason(statuses)
|
||||||
|
if sla_ms > 0 and (time.perf_counter() - start) * 1000 >= sla_ms:
|
||||||
|
return f"KV transfer exceeded SLA {sla_ms}ms"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _collect_transfer_failure_reason(statuses: List[TransferStatus]) -> str:
|
||||||
|
for status in statuses:
|
||||||
|
if status.Failed():
|
||||||
|
return f"KV transfer failed: {status.Message()}"
|
||||||
|
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(
|
||||||
|
self,
|
||||||
|
bootstrap_room: int,
|
||||||
|
kv_indices: npt.NDArray[np.int32],
|
||||||
|
index_slice: slice,
|
||||||
|
is_last_chunk: bool,
|
||||||
|
aux_index: Optional[int] = None,
|
||||||
|
state_indices: Optional[List] = None,
|
||||||
|
num_kv_tokens: Optional[int] = None,
|
||||||
|
wait_event: Optional[object] = None,
|
||||||
|
) -> None:
|
||||||
|
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||||
|
assert not is_last_chunk or (is_last_chunk and aux_index is not None)
|
||||||
|
|
||||||
|
if (
|
||||||
|
bootstrap_room not in self.request_status
|
||||||
|
or self.check_status(bootstrap_room) == KVPoll.Failed
|
||||||
|
):
|
||||||
|
logger.debug(
|
||||||
|
"Request with bootstrap_room=%s already failed", bootstrap_room
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if bootstrap_room not in self.transfer_infos:
|
||||||
|
return
|
||||||
|
|
||||||
|
shard_idx = bootstrap_room % self._num_shards
|
||||||
|
self._transfer_queues[shard_idx].put(
|
||||||
|
TransferKVChunk(
|
||||||
|
room=bootstrap_room,
|
||||||
|
prefill_kv_indices=kv_indices,
|
||||||
|
index_slice=index_slice,
|
||||||
|
is_last_chunk=is_last_chunk,
|
||||||
|
prefill_aux_index=aux_index,
|
||||||
|
state_indices=state_indices,
|
||||||
|
num_kv_tokens=num_kv_tokens,
|
||||||
|
wait_event=wait_event,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def _connect_threadsafe(self, endpoint: str, is_ipv6: bool = False):
|
def _connect_threadsafe(self, endpoint: str, is_ipv6: bool = False):
|
||||||
"""Thread-local ZMQ socket cache with shared Context.
|
"""Thread-local ZMQ socket cache with shared Context.
|
||||||
|
|
||||||
@@ -1298,7 +1451,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
self.kv_args, buffer_index, aux_index, data
|
self.kv_args, buffer_index, aux_index, data
|
||||||
)
|
)
|
||||||
|
|
||||||
def add_transfer_request(
|
def _submit_kv_transfer(
|
||||||
self,
|
self,
|
||||||
bootstrap_room: int,
|
bootstrap_room: int,
|
||||||
kv_indices: npt.NDArray[np.int32],
|
kv_indices: npt.NDArray[np.int32],
|
||||||
@@ -1320,19 +1473,17 @@ class MoriKVManager(CommonKVManager):
|
|||||||
with self.transfer_lock:
|
with self.transfer_lock:
|
||||||
transfer_infos = self.transfer_infos.get(bootstrap_room)
|
transfer_infos = self.transfer_infos.get(bootstrap_room)
|
||||||
if not transfer_infos:
|
if not transfer_infos:
|
||||||
reason = f"No transfer info found for bootstrap_room={bootstrap_room}"
|
raise RuntimeError(
|
||||||
self.record_failure(bootstrap_room, reason)
|
f"No transfer info found for bootstrap_room={bootstrap_room}"
|
||||||
self.update_status(bootstrap_room, KVPoll.Failed)
|
)
|
||||||
return [], None
|
|
||||||
|
|
||||||
self.update_status(bootstrap_room, KVPoll.Transferring)
|
self.update_status(bootstrap_room, KVPoll.Transferring)
|
||||||
for info in transfer_infos.values():
|
for info in transfer_infos.values():
|
||||||
peer_info = self.decode_kv_args_table.get(info.engine_key)
|
peer_info = self.decode_kv_args_table.get(info.engine_key)
|
||||||
if not peer_info:
|
if not peer_info:
|
||||||
reason = f"Peer info missing for engine {info.engine_key}"
|
raise RuntimeError(
|
||||||
self.record_failure(bootstrap_room, reason)
|
f"Peer info missing for engine {info.engine_key}"
|
||||||
self.update_status(bootstrap_room, KVPoll.Failed)
|
)
|
||||||
return [], list(transfer_infos.values())
|
|
||||||
targets.append(TransferTarget(info=info, peer_info=peer_info))
|
targets.append(TransferTarget(info=info, peer_info=peer_info))
|
||||||
if is_last_chunk:
|
if is_last_chunk:
|
||||||
target_infos_snapshot = list(transfer_infos.values())
|
target_infos_snapshot = list(transfer_infos.values())
|
||||||
@@ -1373,15 +1524,11 @@ class MoriKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
reason = f"Transfer submission failed: {e}"
|
|
||||||
with self.transfer_lock:
|
|
||||||
self.record_failure(bootstrap_room, reason)
|
|
||||||
self.update_status(bootstrap_room, KVPoll.Failed)
|
|
||||||
logger.exception(
|
logger.exception(
|
||||||
"Mori KV transfer submission failed for bootstrap_room=%s",
|
"Mori KV transfer submission failed for bootstrap_room=%s",
|
||||||
bootstrap_room,
|
bootstrap_room,
|
||||||
)
|
)
|
||||||
return result_statuses, target_infos_snapshot
|
raise RuntimeError(f"Transfer submission failed: {e}") from e
|
||||||
|
|
||||||
return result_statuses, target_infos_snapshot
|
return result_statuses, target_infos_snapshot
|
||||||
|
|
||||||
@@ -1404,14 +1551,8 @@ class MoriKVSender(CommonKVSender):
|
|||||||
pp_rank,
|
pp_rank,
|
||||||
req_has_disagg_prefill_dp_rank,
|
req_has_disagg_prefill_dp_rank,
|
||||||
)
|
)
|
||||||
self.transfer_statuses: List[TransferStatus] = []
|
|
||||||
self.pending_infos: Optional[List[TransferInfo]] = None
|
|
||||||
self.conclude_state: Optional[KVPoll] = None
|
self.conclude_state: Optional[KVPoll] = None
|
||||||
self.status_notified = False
|
|
||||||
self.init_time = time.time()
|
self.init_time = time.time()
|
||||||
self._notify_lock = threading.Lock()
|
|
||||||
self._notified_status: Optional[KVPoll] = None
|
|
||||||
self._notified_reason: Optional[str] = None
|
|
||||||
|
|
||||||
def send(
|
def send(
|
||||||
self,
|
self,
|
||||||
@@ -1438,199 +1579,57 @@ class MoriKVSender(CommonKVSender):
|
|||||||
self._record_transfer_indices(kv_indices, transfer_state_indices)
|
self._record_transfer_indices(kv_indices, transfer_state_indices)
|
||||||
wait_event = getattr(self, "_early_send_wait_event", None)
|
wait_event = getattr(self, "_early_send_wait_event", None)
|
||||||
self._early_send_wait_event = None
|
self._early_send_wait_event = None
|
||||||
self.kv_mgr.enqueue_transfer(
|
|
||||||
_TransferChunk(
|
if not is_last_chunk:
|
||||||
sender=self,
|
self.kv_mgr.add_transfer_request(
|
||||||
kv_indices=kv_indices,
|
self.bootstrap_room,
|
||||||
index_slice=index_slice,
|
kv_indices,
|
||||||
is_last_chunk=is_last_chunk,
|
index_slice,
|
||||||
aux_index=self.aux_index if is_last_chunk else None,
|
False,
|
||||||
normalized_state=normalized_state,
|
num_kv_tokens=num_kv_tokens,
|
||||||
wait_event=wait_event,
|
wait_event=wait_event,
|
||||||
)
|
)
|
||||||
)
|
else:
|
||||||
self._maybe_finalize_if_room_failed()
|
self.kv_mgr.add_transfer_request(
|
||||||
|
self.bootstrap_room,
|
||||||
def _maybe_finalize_if_room_failed(self) -> None:
|
kv_indices,
|
||||||
if self.conclude_state is not None:
|
index_slice,
|
||||||
return
|
True,
|
||||||
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
|
aux_index=self.aux_index,
|
||||||
self._finalize_failure()
|
state_indices=normalized_state,
|
||||||
|
num_kv_tokens=num_kv_tokens,
|
||||||
def _run_chunk(self, task: _TransferChunk) -> None:
|
wait_event=wait_event,
|
||||||
if self.conclude_state is not None:
|
|
||||||
return
|
|
||||||
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
|
|
||||||
self._finalize_failure()
|
|
||||||
return
|
|
||||||
|
|
||||||
# Wait for the prefill forward that produced these KV pages before
|
|
||||||
# issuing the RDMA read (early-send overlaps that forward).
|
|
||||||
if task.wait_event is not None:
|
|
||||||
task.wait_event.synchronize()
|
|
||||||
|
|
||||||
statuses, infos = self.kv_mgr.add_transfer_request(
|
|
||||||
self.bootstrap_room,
|
|
||||||
task.kv_indices,
|
|
||||||
task.index_slice,
|
|
||||||
task.is_last_chunk,
|
|
||||||
aux_index=task.aux_index,
|
|
||||||
state_indices=task.normalized_state,
|
|
||||||
)
|
|
||||||
self.transfer_statuses.extend(statuses)
|
|
||||||
if infos is not None:
|
|
||||||
self.pending_infos = infos
|
|
||||||
|
|
||||||
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
|
|
||||||
self._finalize_failure()
|
|
||||||
return
|
|
||||||
|
|
||||||
rc = self._wait_chunk(statuses)
|
|
||||||
if self.conclude_state is not None:
|
|
||||||
return
|
|
||||||
if rc != StatusCode.SUCCESS:
|
|
||||||
self._finalize_failure(self._collect_failure_reason())
|
|
||||||
return
|
|
||||||
if task.is_last_chunk:
|
|
||||||
self._notify_decode(KVPoll.Success)
|
|
||||||
with self._notify_lock:
|
|
||||||
if self.conclude_state is None:
|
|
||||||
self.conclude_state = self._notified_status
|
|
||||||
if self._notified_status == KVPoll.Success:
|
|
||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
|
|
||||||
|
|
||||||
def _wait_chunk(self, statuses: List[TransferStatus]) -> StatusCode:
|
|
||||||
if not statuses:
|
|
||||||
return StatusCode.SUCCESS
|
|
||||||
|
|
||||||
start = time.perf_counter()
|
|
||||||
sla_ms = self.kv_mgr._transfer_timeout_ms
|
|
||||||
sla_tripped = False
|
|
||||||
|
|
||||||
while True:
|
|
||||||
rc = self.kv_mgr.engine.wait_all(
|
|
||||||
statuses, timeout_ms=self.kv_mgr._wait_poll_ms
|
|
||||||
)
|
)
|
||||||
if rc != StatusCode.IN_PROGRESS:
|
|
||||||
return rc
|
|
||||||
if (
|
|
||||||
sla_ms > 0
|
|
||||||
and not sla_tripped
|
|
||||||
and (time.perf_counter() - start) * 1000 >= sla_ms
|
|
||||||
):
|
|
||||||
sla_tripped = True
|
|
||||||
self._finalize_failure(f"KV transfer exceeded SLA {sla_ms}ms")
|
|
||||||
|
|
||||||
def _fail_from_worker(self, reason: str) -> None:
|
|
||||||
self._finalize_failure(reason)
|
|
||||||
|
|
||||||
def poll(self) -> KVPoll:
|
def poll(self) -> KVPoll:
|
||||||
if self.conclude_state is not None:
|
if self.conclude_state is not None:
|
||||||
return self.conclude_state
|
return self.conclude_state
|
||||||
|
|
||||||
if self.bootstrap_room not in self.kv_mgr.request_status:
|
if self.bootstrap_room not in self.kv_mgr.request_status:
|
||||||
sent_status, _ = self._finalize_failure()
|
self.conclude_state = KVPoll.Failed
|
||||||
return sent_status
|
return self.conclude_state
|
||||||
|
|
||||||
status = self.kv_mgr.check_status(self.bootstrap_room)
|
status = self.kv_mgr.check_status(self.bootstrap_room)
|
||||||
|
|
||||||
if status == KVPoll.Bootstrapping:
|
if status == KVPoll.Bootstrapping:
|
||||||
elapsed = time.time() - self.init_time
|
timeout_result = self._check_bootstrap_timeout()
|
||||||
if elapsed >= self.kv_mgr.bootstrap_timeout:
|
if timeout_result is not None:
|
||||||
logger.warning_once(
|
self.conclude_state = timeout_result
|
||||||
"Some requests timed out when bootstrapping, "
|
return timeout_result
|
||||||
"which means prefill instances fail to receive the KV indices from the decode instance of this request. "
|
if status in (KVPoll.Success, KVPoll.Failed):
|
||||||
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
|
self.conclude_state = status
|
||||||
)
|
|
||||||
reason = (
|
|
||||||
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s "
|
|
||||||
"in KVPoll.Bootstrapping"
|
|
||||||
)
|
|
||||||
sent_status, _ = self._finalize_failure(reason)
|
|
||||||
return sent_status
|
|
||||||
return status
|
|
||||||
|
|
||||||
if status == KVPoll.Failed:
|
|
||||||
sent_status, _ = self._finalize_failure()
|
|
||||||
return sent_status
|
|
||||||
|
|
||||||
if status == KVPoll.Success:
|
|
||||||
self.conclude_state = KVPoll.Success
|
|
||||||
return KVPoll.Success
|
|
||||||
|
|
||||||
return status
|
return status
|
||||||
|
|
||||||
def _collect_failure_reason(self) -> str:
|
def clear(self) -> None:
|
||||||
for status in self.transfer_statuses:
|
super().clear()
|
||||||
if status.Failed():
|
with self.kv_mgr._room_notify_lock:
|
||||||
return f"KV transfer failed: {status.Message()}"
|
self.kv_mgr._room_status_notified.pop(self.bootstrap_room, None)
|
||||||
return "KV transfer failed due to unknown reason"
|
|
||||||
|
|
||||||
def _terminalize_locked(
|
|
||||||
self,
|
|
||||||
status: KVPoll,
|
|
||||||
reason: Optional[str] = None,
|
|
||||||
) -> Tuple[KVPoll, Optional[str], Optional[List[TransferInfo]]]:
|
|
||||||
if self.status_notified:
|
|
||||||
return self._notified_status, self._notified_reason, None
|
|
||||||
|
|
||||||
if status == KVPoll.Success:
|
|
||||||
with self.kv_mgr.failure_lock:
|
|
||||||
recorded = self.kv_mgr.failure_records.get(self.bootstrap_room)
|
|
||||||
if recorded is not None:
|
|
||||||
status = KVPoll.Failed
|
|
||||||
reason = recorded
|
|
||||||
elif self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
|
|
||||||
status = KVPoll.Failed
|
|
||||||
reason = reason or "request marked Failed before notify"
|
|
||||||
|
|
||||||
if status == KVPoll.Failed:
|
|
||||||
with self.kv_mgr.failure_lock:
|
|
||||||
self.kv_mgr.failure_records.setdefault(
|
|
||||||
self.bootstrap_room, reason or "KV transfer failed"
|
|
||||||
)
|
|
||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
|
|
||||||
|
|
||||||
infos = self.pending_infos
|
|
||||||
if infos is None:
|
|
||||||
with self.kv_mgr.transfer_lock:
|
|
||||||
room_infos = self.kv_mgr.transfer_infos.get(self.bootstrap_room)
|
|
||||||
infos = list(room_infos.values()) if room_infos is not None else None
|
|
||||||
|
|
||||||
self._notified_status = status
|
|
||||||
self._notified_reason = reason
|
|
||||||
self.status_notified = True
|
|
||||||
return status, reason, infos
|
|
||||||
|
|
||||||
def _notify_decode(
|
|
||||||
self, status: KVPoll, failure_reason: Optional[str] = None
|
|
||||||
) -> Tuple[KVPoll, Optional[str]]:
|
|
||||||
with self._notify_lock:
|
|
||||||
emitted_status, emitted_reason, infos = self._terminalize_locked(
|
|
||||||
status, failure_reason
|
|
||||||
)
|
|
||||||
if infos:
|
|
||||||
self.kv_mgr.notify_decode_status(
|
|
||||||
infos, self.bootstrap_room, emitted_status, emitted_reason
|
|
||||||
)
|
|
||||||
return emitted_status, emitted_reason
|
|
||||||
|
|
||||||
def _finalize_failure(
|
|
||||||
self, failure_reason: Optional[str] = None
|
|
||||||
) -> Tuple[KVPoll, Optional[str]]:
|
|
||||||
if failure_reason is None:
|
|
||||||
with self.kv_mgr.failure_lock:
|
|
||||||
failure_reason = self.kv_mgr.failure_records.get(
|
|
||||||
self.bootstrap_room, "KV transfer failed"
|
|
||||||
)
|
|
||||||
sent_status, sent_reason = self._notify_decode(KVPoll.Failed, failure_reason)
|
|
||||||
self.conclude_state = sent_status
|
|
||||||
return sent_status, sent_reason
|
|
||||||
|
|
||||||
def failure_exception(self):
|
def failure_exception(self):
|
||||||
if self.conclude_state is None:
|
if self.conclude_state is None:
|
||||||
self._finalize_failure()
|
self.conclude_state = KVPoll.Failed
|
||||||
|
|
||||||
self.clear()
|
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
|
||||||
@@ -1640,9 +1639,6 @@ class MoriKVSender(CommonKVSender):
|
|||||||
self.bootstrap_room, failure_reason, is_from_another_rank=is_propagated
|
self.bootstrap_room, failure_reason, is_from_another_rank=is_propagated
|
||||||
)
|
)
|
||||||
|
|
||||||
def abort(self):
|
|
||||||
self._finalize_failure("Aborted by AbortReq.")
|
|
||||||
|
|
||||||
|
|
||||||
class MoriKVReceiver(CommonKVReceiver):
|
class MoriKVReceiver(CommonKVReceiver):
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user