From 8ad76415e2db30cdc35d1863304434ab042591fd Mon Sep 17 00:00:00 2001 From: inkcherry Date: Thu, 27 Aug 2026 21:43:13 +0800 Subject: [PATCH] [PD][mori] Align prefill transfer control plane for unified control plane (#36160) --- .../sglang/srt/disaggregation/common/utils.py | 2 + python/sglang/srt/disaggregation/mori/conn.py | 430 +++++++++--------- 2 files changed, 215 insertions(+), 217 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index 8b574904f..c42f6f65c 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -32,6 +32,8 @@ class TransferKVChunk: # Set when the staging worker first counts this chunk toward the per-room # outstanding count; stays set across re-enqueue on a watermark defer. 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: diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 6772c6a7b..6452fc7d3 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -37,6 +37,7 @@ from sglang.srt.disaggregation.common.conn import ( from sglang.srt.disaggregation.common.utils import ( AuxDataCodec, FastQueue, + TransferKVChunk, group_concurrent_contiguous, pack_int_lists, unpack_int_lists, @@ -288,17 +289,6 @@ class TransferTarget: 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): 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._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): threading.Thread( target=self._transfer_worker, @@ -429,34 +421,195 @@ class MoriKVManager(CommonKVManager): return 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: while True: - task = queue.get() + kv_chunk = queue.get() try: - task.sender._run_chunk(task) + self._process_transfer_chunk(kv_chunk) except Exception as exc: failure_reason = f"transfer worker raised: {exc!r}" try: logger.exception( "Mori transfer worker failed for room %s", - task.sender.bootstrap_room, + kv_chunk.room, ) except Exception: pass try: - task.sender._fail_from_worker(failure_reason) + self._conclude_room_failure(kv_chunk.room, failure_reason) except Exception: try: logger.exception( "Mori transfer worker failover failed for room %s", - task.sender.bootstrap_room, + kv_chunk.room, ) except Exception: 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): """Thread-local ZMQ socket cache with shared Context. @@ -1298,7 +1451,7 @@ class MoriKVManager(CommonKVManager): self.kv_args, buffer_index, aux_index, data ) - def add_transfer_request( + def _submit_kv_transfer( self, bootstrap_room: int, kv_indices: npt.NDArray[np.int32], @@ -1320,19 +1473,17 @@ class MoriKVManager(CommonKVManager): with self.transfer_lock: transfer_infos = self.transfer_infos.get(bootstrap_room) if not transfer_infos: - reason = f"No transfer info found for bootstrap_room={bootstrap_room}" - self.record_failure(bootstrap_room, reason) - self.update_status(bootstrap_room, KVPoll.Failed) - return [], None + raise RuntimeError( + f"No transfer info found for bootstrap_room={bootstrap_room}" + ) self.update_status(bootstrap_room, KVPoll.Transferring) for info in transfer_infos.values(): peer_info = self.decode_kv_args_table.get(info.engine_key) if not peer_info: - reason = f"Peer info missing for engine {info.engine_key}" - self.record_failure(bootstrap_room, reason) - self.update_status(bootstrap_room, KVPoll.Failed) - return [], list(transfer_infos.values()) + raise RuntimeError( + f"Peer info missing for engine {info.engine_key}" + ) targets.append(TransferTarget(info=info, peer_info=peer_info)) if is_last_chunk: target_infos_snapshot = list(transfer_infos.values()) @@ -1373,15 +1524,11 @@ class MoriKVManager(CommonKVManager): ) ) 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( "Mori KV transfer submission failed for bootstrap_room=%s", bootstrap_room, ) - return result_statuses, target_infos_snapshot + raise RuntimeError(f"Transfer submission failed: {e}") from e return result_statuses, target_infos_snapshot @@ -1404,14 +1551,8 @@ class MoriKVSender(CommonKVSender): pp_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.status_notified = False self.init_time = time.time() - self._notify_lock = threading.Lock() - self._notified_status: Optional[KVPoll] = None - self._notified_reason: Optional[str] = None def send( self, @@ -1438,199 +1579,57 @@ class MoriKVSender(CommonKVSender): self._record_transfer_indices(kv_indices, transfer_state_indices) wait_event = getattr(self, "_early_send_wait_event", None) self._early_send_wait_event = None - self.kv_mgr.enqueue_transfer( - _TransferChunk( - sender=self, - kv_indices=kv_indices, - index_slice=index_slice, - is_last_chunk=is_last_chunk, - aux_index=self.aux_index if is_last_chunk else None, - normalized_state=normalized_state, + + if not is_last_chunk: + self.kv_mgr.add_transfer_request( + self.bootstrap_room, + kv_indices, + index_slice, + False, + num_kv_tokens=num_kv_tokens, wait_event=wait_event, ) - ) - self._maybe_finalize_if_room_failed() - - def _maybe_finalize_if_room_failed(self) -> None: - if self.conclude_state is not None: - return - if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed: - self._finalize_failure() - - def _run_chunk(self, task: _TransferChunk) -> None: - 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 + else: + self.kv_mgr.add_transfer_request( + self.bootstrap_room, + kv_indices, + index_slice, + True, + aux_index=self.aux_index, + state_indices=normalized_state, + num_kv_tokens=num_kv_tokens, + wait_event=wait_event, ) - 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: if self.conclude_state is not None: return self.conclude_state if self.bootstrap_room not in self.kv_mgr.request_status: - sent_status, _ = self._finalize_failure() - return sent_status + self.conclude_state = KVPoll.Failed + return self.conclude_state status = self.kv_mgr.check_status(self.bootstrap_room) - if status == KVPoll.Bootstrapping: - elapsed = time.time() - self.init_time - if elapsed >= self.kv_mgr.bootstrap_timeout: - logger.warning_once( - "Some requests timed out when bootstrapping, " - "which means prefill instances fail to receive the KV indices from the decode instance of this request. " - "If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT=600' (10 minutes) to relax the timeout condition. " - ) - 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 - + timeout_result = self._check_bootstrap_timeout() + if timeout_result is not None: + self.conclude_state = timeout_result + return timeout_result + if status in (KVPoll.Success, KVPoll.Failed): + self.conclude_state = status return status - def _collect_failure_reason(self) -> str: - for status in self.transfer_statuses: - if status.Failed(): - return f"KV transfer failed: {status.Message()}" - 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 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): if self.conclude_state is None: - self._finalize_failure() + self.conclude_state = KVPoll.Failed + self.clear() + with self.kv_mgr.failure_lock: failure_reason = self.kv_mgr.failure_records.pop(self.bootstrap_room, 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 ) - def abort(self): - self._finalize_failure("Aborted by AbortReq.") - class MoriKVReceiver(CommonKVReceiver):