diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index c8f90f987..b175ca541 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -296,6 +296,7 @@ class _TransferChunk: 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): @@ -1429,6 +1430,8 @@ class MoriKVSender(CommonKVSender): else None ) self._record_transfer_indices(kv_indices, 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, @@ -1437,6 +1440,7 @@ class MoriKVSender(CommonKVSender): is_last_chunk=is_last_chunk, aux_index=self.aux_index if is_last_chunk else None, normalized_state=normalized_state, + wait_event=wait_event, ) ) self._maybe_finalize_if_room_failed() @@ -1454,6 +1458,11 @@ class MoriKVSender(CommonKVSender): 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, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 43efa2db9..6e8bc0db0 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -1038,6 +1038,15 @@ class SchedulerDisaggregationPrefillMixin: if cached_end <= req.start_send_idx: return assert cached_end % self.token_to_kv_pool_allocator.page_size == 0 + # Early-send issues the KV read before this step's forward is enqueued, + # but under overlap scheduling the PRIOR step's prefill forward may still + # be writing these prefix pages on forward_stream. Record a completion + # event now so the transfer worker can wait on those writes before the + # RDMA read, instead of racing them. + if self.enable_overlap: + ev = torch.cuda.Event() + ev.record(self.forward_stream) + req.disagg_kv_sender._early_send_wait_event = ev self.send_kv_chunk(req, last_chunk=False, end_idx=cached_end) def send_kv_chunk(