Per code review (HiCache load is layer-by-layer & correctly ordered into the compute stream before the backup hook; current-reuse is a within-forward read that doesn't rewrite pool pages): unify registration to the EXACT range this forward's send_kv_chunk transmits — req_to_token[start_send_idx:end_idx], page-floored for a non-last chunk. Non-chunked = one full range; chunked = one range per chunk. Drop the is_chunked/start_send_idx skip. To avoid the review's collision risk (chunk N still finishing when chunk N+1 registers), the manager keys contexts per (room, start_send_idx): _active[room] is a FIFO deque of (chunk_key, ctx); register dedups the same chunk but appends a new one; finish(room) pops the FRONT (chunks finish in send order — no chunk key needed in the transfer_worker); drop drains all the room's chunks; on_layer_end enqueues for all active chunk contexts (per-ctx note_enqueued dedup keeps each chunk's own events). 28 unit tests pass incl. chunked FIFO + per-chunk dedup. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
314 lines
11 KiB
Python
314 lines
11 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_repeated_enqueues_dedup_so_finish_cannot_hang(self):
|
|
# the notifier re-fires every layer on every subsequent forward while the ctx
|
|
# is still active; note_enqueued must dedup so target stays capped at the layer
|
|
# count and finish() completes (the high-cache-hit hang was target>>layers).
|
|
eng = _FakeEngine()
|
|
ctx = _ctx(eng, num_layers=3)
|
|
for _forward in range(5): # 5 forwards all re-fire layers 0,1,2
|
|
for L in range(3):
|
|
self.assertEqual(ctx.note_enqueued(L), _forward == 0) # True only 1st
|
|
if _forward == 0:
|
|
ctx.submit_layer(L, None)
|
|
self.assertEqual(len(eng.submits), 3) # 3 unique despite 15 note_enqueued
|
|
self.assertEqual(ctx.finish(timeout=0.5), 0) # completes, no 30s timeout
|
|
|
|
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 -> all 3 moved, success. The fallback BATCHES the missing
|
|
# layers into one submit.
|
|
eng = _FakeEngine()
|
|
ctx = _ctx(eng, num_layers=3)
|
|
ctx.submit_layer(0, None) # notifier fired layer 0 only -> 1 submit
|
|
self.assertEqual(ctx.finish(timeout=0.5), 0)
|
|
self.assertEqual(len(eng.submits), 2) # layer 0 + fallback batch [1,2]
|
|
self.assertEqual(eng.waits[-1], [101, 102])
|
|
self.assertEqual(eng.submits[1], ("sess", [1000, 2000], [6000, 7000], [64, 64]))
|
|
|
|
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:
|
|
num_layers = 8
|
|
|
|
def __init__(self):
|
|
self.submitted = [] # (start, end, event)
|
|
self.enqueued = 0
|
|
self.finished = False
|
|
self.failed = False
|
|
self._enq = set()
|
|
|
|
def note_enqueued(self, key):
|
|
if key in self._enq:
|
|
return False
|
|
self._enq.add(key)
|
|
self.enqueued += 1
|
|
return True
|
|
|
|
def submit_group(self, start, end, event):
|
|
self.submitted.append((start, end, 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, group_size=1):
|
|
from sglang.srt.disaggregation.cp_per_layer_transfer import PerLayerTransferManager
|
|
|
|
return PerLayerTransferManager(
|
|
num_workers=num_workers,
|
|
event_factory=_RecEvent,
|
|
current_stream=lambda: "STREAM",
|
|
group_size=group_size,
|
|
)
|
|
|
|
|
|
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(group_size=1) # per-layer: every layer is a group boundary
|
|
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, start, end, ev in items:
|
|
self.assertEqual((start, end), (5, 5))
|
|
self.assertEqual(ev.recorded_stream, "STREAM")
|
|
|
|
def test_grouping_only_enqueues_at_boundaries(self):
|
|
m = _manager(group_size=4)
|
|
c = _MockCtx()
|
|
m.register("r1", c)
|
|
for L in range(8):
|
|
m.on_layer_end(L)
|
|
items = _drain(m._q)
|
|
# boundaries at L=3 (group [0,3]) and L=7 (group [4,7]) -> 2 groups, not 8
|
|
self.assertEqual([(s, e) for _, s, e, _ in items], [(0, 3), (4, 7)])
|
|
|
|
def test_chunked_fifo_separate_contexts_and_dedup(self):
|
|
m = _manager(group_size=1)
|
|
c1, c2 = _MockCtx(), _MockCtx()
|
|
m.register("r1", c1, chunk_key=0) # chunk 1
|
|
m.register("r1", c1, chunk_key=0) # SAME chunk re-register -> deduped
|
|
m.register("r1", c2, chunk_key=64) # chunk 2 (different start_send_idx)
|
|
self.assertTrue(m.has_chunk("r1", 0))
|
|
self.assertTrue(m.has_chunk("r1", 64))
|
|
self.assertFalse(m.has_chunk("r1", 128))
|
|
m.on_layer_end(5) # both chunk contexts active -> both enqueued
|
|
self.assertEqual(len(_drain(m._q)), 2)
|
|
# finish pops FIFO: chunk 1 (c1) first, then chunk 2 (c2)
|
|
self.assertEqual(m.finish("r1"), 0)
|
|
self.assertTrue(c1.finished)
|
|
self.assertFalse(c2.finished)
|
|
self.assertEqual(m.finish("r1"), 0)
|
|
self.assertTrue(c2.finished)
|
|
self.assertFalse(m.has_room("r1"))
|
|
|
|
def test_worker_step_calls_submit_group(self):
|
|
m = _manager()
|
|
c = _MockCtx()
|
|
ev = _RecEvent()
|
|
m._worker_step((c, 0, 3, ev))
|
|
self.assertEqual(c.submitted, [(0, 3, ev)])
|
|
|
|
def test_worker_step_marks_failed_on_exception(self):
|
|
m = _manager()
|
|
|
|
class _Boom:
|
|
def __init__(self):
|
|
self.failed = False
|
|
|
|
def submit_group(self, start, end, event):
|
|
raise RuntimeError("boom")
|
|
|
|
def mark_failed(self):
|
|
self.failed = True
|
|
|
|
b = _Boom()
|
|
m._worker_step((b, 0, 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(group_size=1)
|
|
m.on_layer_end(0)
|
|
self.assertEqual(_drain(m._q), [])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|