From 544cdd0471460c5d656895718cbe8442889f1268 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E6=97=AD?= Date: Sat, 5 Sep 2026 11:58:58 +0800 Subject: [PATCH] fix(nixl): prevent early staging buffer reuse and duplicate sends MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: 刘旭 --- .../sglang/srt/disaggregation/common/utils.py | 4 + python/sglang/srt/disaggregation/nixl/conn.py | 51 ++++--- .../disaggregation/test_nixl_backend_basic.py | 131 ++++++++++++++++++ 3 files changed, 170 insertions(+), 16 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index 1571c3d19372..0c7c8a126d32 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -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: diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 8e1ceb7fd00f..0df9a0a57b66 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -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). @@ -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( @@ -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] @@ -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] @@ -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 @@ -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. diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index c40cc8b1e166..280989d78b05 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -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):