Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion python/sglang/srt/disaggregation/common/staging_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,8 +237,17 @@ def free(self, alloc_id: int):
self.alloc_order.pop(0)

if not self.allocations:
# Once the ring is empty, every byte is safe to reuse. Keeping
# ``head`` at the end of the last allocation leaves the unused
# tail represented as unsafe after the next wrap: an allocation
# in the new round whose end extends past the old head can then
# wait forever for a watermark that can no longer advance.
# Start a fresh round at offset zero so the empty-ring
# watermark covers the whole previous round.
self.round += 1
self.head = 0
self.watermark_round = self.round
self.watermark_tail = self.head
self.watermark_tail = 0
elif self.alloc_order:
off, _, rnd = self.allocations[self.alloc_order[0]]
self.watermark_round = rnd
Expand Down
47 changes: 31 additions & 16 deletions python/sglang/srt/disaggregation/common/staging_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,8 +107,21 @@ def register_wm_subscriber(self, receiver, session_id: str) -> None:
if receiver is None or not receiver.bootstrap_infos:
return
key = tuple(str(bi) for bi in receiver.bootstrap_infos)
if key not in self._wm_subscribers:
self._wm_subscribers[key] = (receiver, session_id)
if key in self._wm_subscribers:
return

self._wm_subscribers[key] = (receiver, session_id)
# The allocator round is global, while prefills learn watermarks per
# session. A newly registered prefill may therefore receive its first
# allocation after another session has already advanced the ring. Send
# the current watermark immediately; waiting for the next free can
# deadlock when that first allocation is itself waiting on the missed
# watermark.
self._send_watermark(
receiver,
session_id,
self.staging_allocator.get_watermark(),
)

def num_writers_for(self, receiver) -> int:
"""Compute all TP and PP writers expected for a staging chunk."""
Expand Down Expand Up @@ -442,26 +455,28 @@ def _scatter_region(
return True

def _free_and_send_watermark(
self, alloc_id: int, decode_req: DecodeRequest
self, alloc_id: int, _decode_req: DecodeRequest
) -> None:
"""Free a staging allocation and broadcast watermark to all prefills."""
self.staging_allocator.free(alloc_id)
post_wm = self.staging_allocator.get_watermark()
room = decode_req.req.bootstrap_room
wm_round, wm_tail = post_wm
for receiver, session_id in list(self._wm_subscribers.values()):
self._send_watermark(receiver, session_id, post_wm)

@staticmethod
def _send_watermark(receiver, session_id: str, watermark) -> None:
"""Send one allocator watermark to a registered prefill session."""
wm_round, wm_tail = watermark
wm_round_b = str(wm_round).encode("ascii")
wm_tail_b = str(wm_tail).encode("ascii")
for _key, (receiver, session_id) in list(self._wm_subscribers.items()):
sid_b = session_id.encode("ascii")
for bootstrap_info in receiver.bootstrap_infos:
try:
sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info)
with lock:
sock.send_multipart(
[b"WATERMARK", wm_round_b, wm_tail_b, sid_b]
)
except Exception:
pass
sid_b = session_id.encode("ascii")
for bootstrap_info in receiver.bootstrap_infos:
try:
sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info)
with lock:
sock.send_multipart([b"WATERMARK", wm_round_b, wm_tail_b, sid_b])
except Exception:
pass


def is_watermark_ready(
Expand Down
6 changes: 6 additions & 0 deletions python/sglang/srt/disaggregation/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,12 @@ 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
# A heterogeneous-TP chunk fans out to multiple decode sessions. Their
# staging watermarks advance independently, so one worker pass can finish
# only a prefix of the fan-out before another destination defers. Keep the
# completed sessions on the re-enqueued chunk: replaying them can write into
# a staging slot that decode has already scattered and recycled.
staging_completed_sessions: set[str] = dataclasses.field(default_factory=set)
# Mori early-send: CUDA event to synchronize before RDMA (optional).
wait_event: Optional[object] = None

Expand Down
Loading
Loading