diff --git a/python/sglang/srt/disaggregation/cp_per_layer_transfer.py b/python/sglang/srt/disaggregation/cp_per_layer_transfer.py new file mode 100644 index 000000000..54abe49e8 --- /dev/null +++ b/python/sglang/srt/disaggregation/cp_per_layer_transfer.py @@ -0,0 +1,82 @@ +"""Per-layer overlapped KV transfer (lever A) — per-request coordinator. + +Design (see docs_internal/per-layer-async-transfer-design.md + lever-a-implementation-plan.md): +- The prefill forward fires a per-layer end hook (notify_layer_end_for_backup -> + layer_backup_notifiers) after each layer's KV write. We register a notifier that, + on the forward thread, records a CUDA event E_L (cheap, non-blocking) and enqueues + (context, layer_id, E_L) to a background worker. +- The background worker `event.synchronize()`s E_L (waits the write KERNEL to finish + on the GPU — on the worker thread, NOT the compute stream) and only THEN submits + layer L's RDMA transfer asynchronously (batch_transfer_async, the safe G1 path). + Without the event wait the RDMA could read layer L before the write kernel finished. +- After the forward, finish() waits all submitted batch_ids and reports status. + +This module holds the per-request state + the submit/finish logic; the scheduler +wiring (setup before forward, notifier registration, finish after) lives in the +prefill/conn integration. Engine + events are injected so this is unit-testable. +""" +from __future__ import annotations + +import threading +from typing import Callable, List, Optional, Tuple + +# get_blocks(layer_id) -> (src_addrs, dst_addrs, lengths) for THIS rank's owned pages +# of layer_id, or None to skip (no owned pages / dummy). +BlocksFn = Callable[[int], Optional[Tuple[List[int], List[int], List[int]]]] + + +class PerLayerTransferContext: + """Tracks the per-layer async transfers for one prefill request.""" + + def __init__(self, engine, session_id: str, get_blocks: BlocksFn): + self.engine = engine + self.session_id = session_id + self.get_blocks = get_blocks + self._batch_ids: List[int] = [] + self._submitted_layers: set[int] = set() + self._failed = False + self._lock = threading.Lock() + + @property + def failed(self) -> bool: + with self._lock: + return self._failed + + def submit_layer(self, layer_id: int, event) -> None: + """Wait the layer-L write event, then async-submit layer L's transfer. + Called on a BACKGROUND worker thread (event.synchronize() must not run on + the compute/forward thread). Idempotent per layer; no-op after failure.""" + with self._lock: + if self._failed or layer_id in self._submitted_layers: + return + self._submitted_layers.add(layer_id) + if event is not None: + event.synchronize() # wait the GPU write kernel for layer_id to finish + blocks = self.get_blocks(layer_id) + if not blocks: + return + src_addrs, dst_addrs, lengths = blocks + if not src_addrs: + return + batch_id = self.engine.batch_transfer_async_submit( + self.session_id, list(src_addrs), list(dst_addrs), list(lengths) + ) + with self._lock: + if batch_id < 0: + self._failed = True + else: + self._batch_ids.append(batch_id) + + def finish(self) -> int: + """Wait for all submitted layer transfers. Returns 0 on success, non-zero on + failure (a submit failed, or wait reported failure). Drains regardless.""" + with self._lock: + failed = self._failed + batch_ids = list(self._batch_ids) + wait_status = self.engine.wait_batch_transfers(batch_ids) + return -1 if failed else wait_status + + @property + def num_submitted(self) -> int: + with self._lock: + return len(self._batch_ids) diff --git a/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py b/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py new file mode 100644 index 000000000..afe5e97c1 --- /dev/null +++ b/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py @@ -0,0 +1,107 @@ +"""Unit tests for A2 core: PerLayerTransferContext (per-request overlapped transfer). + +Key correctness properties: +- submit_layer waits the layer's CUDA write event BEFORE issuing the RDMA (so the + transfer never reads a layer before its write kernel finished); +- idempotent per layer; skips layers with no owned blocks; +- a submit failure marks the context failed and stops further submits; +- finish waits all accumulated batch_ids and returns 0 / non-zero correctly. + +Engine and CUDA events are mocked — no GPU/RDMA. +""" +import unittest + +from sglang.srt.disaggregation.cp_per_layer_transfer import PerLayerTransferContext + + +class _FakeEvent: + def __init__(self): + self.synced = False + + def synchronize(self): + self.synced = True + + +class _FakeEngine: + def __init__(self, submit_rets=None, wait_ret=0): + self.submits = [] + self.waits = [] + self._rets = submit_rets + self._wait = wait_ret + self._n = 0 + self.submit_order = [] # (layer-derived addr, whether event was synced first) + + def batch_transfer_async_submit(self, sid, src, dst, lens): + self.submits.append((sid, list(src), list(dst), list(lens))) + if self._rets is not None: + r = self._rets[self._n] + self._n += 1 + return r + return 100 + len(self.submits) + + def wait_batch_transfers(self, ids): + self.waits.append(list(ids)) + return self._wait + + +def _blocks(layer_id): + base = 1000 * layer_id + return ([base], [base + 5000], [64]) + + +def _ctx(engine, blocks=_blocks): + return PerLayerTransferContext(engine, "sess", blocks) + + +class TestPerLayerTransferContext(unittest.TestCase): + def test_submit_waits_event_then_submits(self): + eng = _FakeEngine() + ev = _FakeEvent() + _ctx(eng).submit_layer(3, ev) + self.assertTrue(ev.synced) # write event waited BEFORE the RDMA submit + self.assertEqual(eng.submits, [("sess", [3000], [8000], [64])]) + + def test_idempotent_per_layer(self): + eng = _FakeEngine() + ctx = _ctx(eng) + ctx.submit_layer(0, _FakeEvent()) + ctx.submit_layer(0, _FakeEvent()) + self.assertEqual(len(eng.submits), 1) + + def test_skip_layer_with_no_blocks(self): + eng = _FakeEngine() + + def blocks(L): + return None if L == 1 else ([L], [L + 1], [8]) + + ctx = _ctx(eng, blocks) + for L in range(3): + ctx.submit_layer(L, None) + self.assertEqual(len(eng.submits), 2) # layer 1 (no blocks) skipped + + def test_finish_waits_all_and_returns_zero(self): + eng = _FakeEngine() + ctx = _ctx(eng) + for L in range(3): + ctx.submit_layer(L, _FakeEvent()) + self.assertEqual(ctx.finish(), 0) + self.assertEqual(eng.waits, [[101, 102, 103]]) + + def test_submit_failure_marks_failed_and_stops(self): + eng = _FakeEngine(submit_rets=[101, -1, 103]) + ctx = _ctx(eng) + for L in range(3): + ctx.submit_layer(L, None) + self.assertTrue(ctx.failed) + self.assertEqual(len(eng.submits), 2) # stopped after the -1 + self.assertEqual(ctx.finish(), -1) + + def test_finish_propagates_wait_failure(self): + eng = _FakeEngine(wait_ret=-9) + ctx = _ctx(eng) + ctx.submit_layer(0, None) + self.assertEqual(ctx.finish(), -9) + + +if __name__ == "__main__": + unittest.main()