Add PerLayerTransferContext: the per-request coordinator for overlapped per-layer KV transfer. submit_layer(layer_id, event) waits the layer's CUDA write event (on a background thread, never the compute stream) before async-submitting that layer's RDMA via the G1 path, so the transfer never reads a layer before its write kernel finished — the core correctness invariant for the forward overlap. finish() waits all accumulated batch_ids; idempotent per layer; fails closed. Unit-tested (test_cp_per_layer_transfer.py, 6 cases): event-wait-before-submit ordering, idempotency, empty-layer skip, finish-waits-all, submit-failure stop, wait-failure propagation. The scheduler/notifier wiring (A2-wiring + A3) builds on this. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
108 lines
3.3 KiB
Python
108 lines
3.3 KiB
Python
"""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()
|