Files
sglang/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py
leavelet 64f767bb42 fix(disagg): cover notifier-missed layers via post-forward fallback (lever A)
E2E diagnostic was precise: "finish TIMEOUT processed=78/79", submit_failed=0 —
the per-layer notifier fires 78x but kv_data_ptrs has 79 layers (the 79th is the
MTP/nextn EAGLE buffer: present in kv_data_ptrs so the monolithic send moves it,
but it doesn't fire the per-layer hook). The old completion required all num_layers,
so it both hung to the timeout AND would silently miss that layer's KV (corruption).

Redesign: gate completion on the ACTUAL enqueued count (note_enqueued), and in
finish() SYNCHRONOUSLY transfer any layers the notifier didn't fire for (KV is fully
written post-forward, no event needed). Net: the per-layer set == kv_data_ptrs,
byte-identical to the monolithic send; robust to any model firing fewer hooks than
KV buffers. The fired layers stay overlapped with the forward.

Unit tests updated (25 pass): fallback transfers the missed layers; submit failures
still report -1.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-07 09:51:05 +00:00

259 lines
8.2 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, num_layers=8):
return PerLayerTransferContext(engine, "sess", blocks, num_layers=num_layers)
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, num_layers=3)
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, num_layers=3)
for L in range(3):
ctx.submit_layer(L, None)
self.assertTrue(ctx.failed)
self.assertEqual(len(eng.submits), 2) # stopped submitting after the -1
self.assertEqual(ctx.finish(), -1)
def test_finish_propagates_wait_failure(self):
eng = _FakeEngine(wait_ret=-9)
ctx = _ctx(eng, num_layers=1)
ctx.submit_layer(0, None)
self.assertEqual(ctx.finish(), -9)
def test_finish_falls_back_for_notifier_missed_layers(self):
# only layer 0 fired by the notifier; finish must SYNCHRONOUSLY transfer the
# missing layers 1,2 (e.g. an MTP buffer not hooked) -> all 3 moved, success.
eng = _FakeEngine()
ctx = _ctx(eng, num_layers=3)
ctx.submit_layer(0, None) # notifier fired layer 0 only
self.assertEqual(ctx.finish(timeout=0.5), 0)
self.assertEqual(len(eng.submits), 3) # layer 0 (overlap) + 1,2 (fallback)
self.assertEqual(eng.waits[-1], [101, 102, 103])
def test_failed_submit_still_reported_after_fallback(self):
# layer 2 submit fails during the overlap; finish reports -1 even though it
# also runs the fallback for any unfired layers.
eng = _FakeEngine(submit_rets=[101, -1, 103])
ctx = _ctx(eng, num_layers=3)
for L in range(3):
ctx.submit_layer(L, None)
self.assertEqual(ctx.finish(timeout=0.5), -1)
class TestBuildLayerBlocks(unittest.TestCase):
def test_addresses_and_lengths(self):
from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks
src, dst, lens = build_layer_blocks(
1000, 5000, 64, [[0, 1, 2], [5, 6]], [[10, 11, 12], [20, 21]]
)
self.assertEqual(src, [1000 + 0 * 64, 1000 + 5 * 64])
self.assertEqual(dst, [5000 + 10 * 64, 5000 + 20 * 64])
self.assertEqual(lens, [64 * 3, 64 * 2]) # item_len * run length
def test_empty(self):
from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks
self.assertEqual(build_layer_blocks(1, 2, 8, [], []), ([], [], []))
def test_only_base_ptr_changes_across_layers(self):
from sglang.srt.disaggregation.cp_per_layer_transfer import build_layer_blocks
s0, d0, l0 = build_layer_blocks(1000, 2000, 64, [[3]], [[7]])
s1, d1, l1 = build_layer_blocks(9000, 8000, 64, [[3]], [[7]])
self.assertEqual(s0, [1000 + 3 * 64])
self.assertEqual(s1, [9000 + 3 * 64])
self.assertEqual(l0, l1) # lengths identical across layers (the invariant)
class _MockCtx:
def __init__(self):
self.submitted = []
self.enqueued = 0
self.finished = False
self.failed = False
def note_enqueued(self):
self.enqueued += 1
def submit_layer(self, layer_id, event):
self.submitted.append((layer_id, event))
def finish(self):
self.finished = True
return 0
def mark_failed(self):
self.failed = True
class _RecEvent:
def __init__(self):
self.recorded_stream = "UNSET"
self.synced = False
def record(self, stream=None):
self.recorded_stream = stream
def synchronize(self):
self.synced = True
def _manager(num_workers=0):
from sglang.srt.disaggregation.cp_per_layer_transfer import PerLayerTransferManager
return PerLayerTransferManager(
num_workers=num_workers, event_factory=_RecEvent, current_stream=lambda: "STREAM"
)
def _drain(q):
import queue as _q
items = []
while True:
try:
items.append(q.get(block=False))
except _q.Empty:
break
return items
class TestPerLayerTransferManager(unittest.TestCase):
def test_on_layer_end_enqueues_per_active_ctx_with_recorded_event(self):
m = _manager()
c1, c2 = _MockCtx(), _MockCtx()
m.register("r1", c1)
m.register("r2", c2)
m.on_layer_end(5)
items = _drain(m._q)
self.assertEqual(len(items), 2) # one per active context
for ctx, layer_id, ev in items:
self.assertEqual(layer_id, 5)
self.assertEqual(ev.recorded_stream, "STREAM") # event recorded on compute stream
def test_worker_step_calls_submit_layer(self):
m = _manager()
c = _MockCtx()
ev = _RecEvent()
m._worker_step((c, 3, ev))
self.assertEqual(c.submitted, [(3, ev)])
def test_worker_step_marks_failed_on_exception(self):
m = _manager()
class _Boom:
def __init__(self):
self.failed = False
def submit_layer(self, layer_id, event):
raise RuntimeError("boom")
def mark_failed(self):
self.failed = True
b = _Boom()
m._worker_step((b, 0, None))
self.assertTrue(b.failed)
def test_finish_pops_and_calls_ctx_finish(self):
m = _manager()
c = _MockCtx()
m.register("r1", c)
self.assertTrue(m.has_active())
self.assertEqual(m.finish("r1"), 0)
self.assertTrue(c.finished)
self.assertFalse(m.has_active())
self.assertEqual(m.finish("r1"), 0) # idempotent after pop
def test_on_layer_end_no_active_is_noop(self):
m = _manager()
m.on_layer_end(0)
self.assertEqual(_drain(m._q), [])
if __name__ == "__main__":
unittest.main()