Skip to content
Open
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
4 changes: 4 additions & 0 deletions python/sglang/srt/disaggregation/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,10 @@ class TransferKVChunk:
staging_counted: bool = False
# Mori early-send: CUDA event to synchronize before RDMA (optional).
wait_event: Optional[object] = None
# Protocol progress belongs to the work item, so a staging deferral can
# retry unfinished destinations without replaying completed KV/aux/state.
# Values are destination session identities, never transport handles.
completed_destinations: set[str] = dataclasses.field(default_factory=set)


def pack_list_of_buffers(buffers: List[bytes]) -> bytes:
Expand Down
51 changes: 35 additions & 16 deletions python/sglang/srt/disaggregation/nixl/conn.py
Original file line number Diff line number Diff line change
Expand Up @@ -1083,6 +1083,20 @@ def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
dst_mem_kind=dst_mem_kind,
)

def _wait_for_transfer_handles(self, handles: List[Any], room: int) -> None:
"""Wait for the submitted group using the existing NIXL error policy."""
while handles:
all_done = True
for handle in handles:
state = self.agent.check_xfer_state(handle)
if state == "ERR":
raise RuntimeError(f"NIXL transfer encountered ERR room={room}")
if state != "DONE":
all_done = False
if all_done:
return
time.sleep(0)

def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0):
# Per-worker staging strategy: lazy-created on first chunk so we
# see kv_buffer_tensors (set by ModelRunner after engine init).
Expand All @@ -1093,6 +1107,7 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)
kv_chunk: TransferKVChunk = queue.get()
room = kv_chunk.room
handles: List[Any] = []
submitted_destinations = set()
try:
if room not in self.request_status:
logger.debug(
Expand Down Expand Up @@ -1154,6 +1169,8 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)
assert room == req.room
if req.is_dummy():
continue
if req.agent_name in kv_chunk.completed_destinations:
continue

assert req.agent_name in self.decode_kv_args_table
dst_info = self.decode_kv_args_table[req.agent_name]
Expand Down Expand Up @@ -1304,6 +1321,11 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)

if kv_xfer_handle is not None:
handles.append(kv_xfer_handle)
if use_staging:
# The next destination gathers into this same
# worker buffer. Its previous asynchronous
# reader must finish before that overwrite.
self._wait_for_transfer_handles([kv_xfer_handle], room)

if kv_chunk.is_last_chunk:
dst_info = self.decode_kv_args_table[req.agent_name]
Expand Down Expand Up @@ -1347,24 +1369,21 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0)
)
handles.append(aux_xfer_handle)

submitted_destinations.add(req.agent_name)

# A partial fanout may already have posted KV, aux or state
# before another destination defers. Drain those handles even
# on requeue, then preserve the completed destinations on the
# common work item. This also protects DCP pack regions from
# being reused by a different room while these reads are live.
self._wait_for_transfer_handles(handles, room)
if self.enable_staging:
kv_chunk.completed_destinations.update(submitted_destinations)

if staging_deferred:
# Chunk has been re-enqueued; do not advance status.
continue

while handles:
all_done = True
for handle in handles:
state = self.agent.check_xfer_state(handle)
if state == "ERR":
raise RuntimeError(
f"NIXL transfer encountered ERR room={room}"
)
if state != "DONE":
all_done = False
if all_done:
break
time.sleep(0)

self._staging_outstanding[room] -= 1
if self.enable_deferred_decode_kv_release:
# Handles all DONE => this room's writes landed; ack if it
Expand Down Expand Up @@ -1986,8 +2005,8 @@ def _do_staging_transfer(
retried on the next pop.
- oversized chunk (will never fit) -> raise RuntimeError.
- staging successfully posted -> return ``(handle, False)``. The
caller appends the handle to the per-chunk handle list and
busy-polls it to DONE alongside other handles.
caller waits for this handle before the next gather can reuse
the worker's source buffer.
- send_kvcache_staged returned None (chunk cannot fit; decode buffer
too small, kv_buffer_tensors missing, etc.) -> raise RuntimeError
instead of falling back to the slice path.
Expand Down
131 changes: 131 additions & 0 deletions test/registered/unit/disaggregation/test_nixl_backend_basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -621,6 +621,137 @@ def check_xfer_state(_handle):
)
self.assertEqual(submitted_counts_at_poll, [3, 3, 3])

def _run_staging_fanout(self, *, defer_second=False, mixed_first=False):
room = 24
mgr = self._make_manager(room)
mgr.enable_staging = True
mgr._staging_ctx = PrefillStagingContext()
template_req = mgr.transfer_infos[room]["agent"]
template_dst = mgr.decode_kv_args_table["agent"]
mgr.transfer_infos[room] = {}
mgr.decode_kv_args_table = {}
for rank, peer in enumerate(("a", "b")):
mgr.transfer_infos[room][peer] = TransferInfo(
**{**vars(template_req), "agent_name": peer}
)
mgr.decode_kv_args_table[peer] = SimpleNamespace(
**{
**vars(template_dst),
"decode_tp_size": 2,
"decode_tp_rank": rank,
"staging_base_ptr": 0x1000,
"staging_total_size": 4096,
"dst_kv_item_len": 4,
"dst_state_data_ptrs": [0],
"dst_state_item_lens": [4],
"dst_state_dim_per_tensor": [1],
"dst_state_layer_ids": [0],
}
)
if mixed_first:
mgr.decode_kv_args_table["a"].decode_tp_size = 1
mgr.decode_kv_args_table["a"].kv_xfer_segments = [object()]

chunk = self._make_chunk(room, [1], is_last_chunk=True)
chunk.state_indices = [0]
pending = [chunk]
handles = []
reads = []
live_at_dequeue = []
source = [None]

def get():
live_at_dequeue.append([h for h in handles if not h["done"]])
if not pending:
raise SystemExit()
return pending.pop(0)

def post(peer, kind):
handle = dict(peer=peer, kind=kind, polls=0, done=False)
handles.append(handle)
return handle

def staged(peer, *args, **kwargs):
# Model a destination-specific gather into the worker's one buffer.
# NIXL reads it asynchronously, when this handle completes below.
source[0] = peer
return post(peer, "staged")

def check(handle):
if handle["done"]:
return "DONE"
handle["polls"] += 1
if handle["polls"] == 1:
return "PROC"
if handle["kind"] == "staged":
reads.append((handle["peer"], source[0]))
handle["done"] = True
return "DONE"

ready_calls = defaultdict(int)

def ready(req, *args, **kwargs):
ready_calls[req.agent_name] += 1
if defer_second and req.agent_name == "b" and ready_calls["b"] == 1:
return False, 0, -1, 0, -1
return True, 0, 0, 0, 0

strategy = SimpleNamespace(
check_ready=ready, staging_buffer=FakeStagingBuffer()
)
mgr._try_create_staging_strategy = lambda buffer: strategy
mgr.send_kvcache_staged = staged
mgr.send_kvcache_mixed = lambda peer, *args: [
post(peer, "mixed-vram"),
post(peer, "mixed-dram"),
]
mgr.send_aux = lambda peer, *args: post(peer, "aux")
mgr.maybe_send_extra = lambda peer, *args, **kwargs: [post(peer, "state")]
mgr.agent = SimpleNamespace(check_xfer_state=check)
queue = SimpleNamespace(get=get, put=pending.append)
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.common.staging_buffer": _fake_staging_buffer_module()
},
):
with self.assertRaises(SystemExit):
mgr.transfer_worker(queue, staging_buffer=strategy.staging_buffer)

self.assertEqual(mgr.exceptions, {})
self.assertEqual(mgr.request_status[room], KVPoll.Success)
self.assertEqual(mgr._staging_outstanding.get(room, 0), 0)
return handles, reads, live_at_dequeue

def test_staging_fanout_preserves_each_destinations_source(self):
_, reads, _ = self._run_staging_fanout()
self.assertEqual(reads, [("a", "a"), ("b", "b")])

def test_staging_fanout_deferral_does_not_replay_completed_destination(self):
handles, reads, live = self._run_staging_fanout(defer_second=True)
self.assertEqual(
[(h["peer"], h["kind"]) for h in handles],
[
(peer, kind)
for peer in ("a", "b")
for kind in ("staged", "state", "aux")
],
)
self.assertEqual(reads, [("a", "a"), ("b", "b")])
self.assertTrue(all(not pending for pending in live))

def test_staging_fanout_deferral_drains_mixed_handles_before_next_dequeue(self):
handles, reads, live = self._run_staging_fanout(
defer_second=True, mixed_first=True
)
self.assertTrue(all(not pending for pending in live))
self.assertEqual(
[(h["peer"], h["kind"]) for h in handles],
[("a", kind) for kind in ("mixed-vram", "mixed-dram", "state", "aux")]
+ [("b", kind) for kind in ("staged", "state", "aux")],
)
self.assertEqual(reads, [("b", "b")])


class TestNixlNotifications(CustomTestCase):
def _make_manager(self, messages, required=None):
Expand Down
Loading