perf(disagg): coarse-grained per-layer transfer + skip re-registration (lever A)

For large bs (production target: bs~10 x ~100k tokens), per-layer granularity does
10x79 = 790 submitTransfer calls + CUDA events + enqueues per forward on the forward
thread. Two overhead cuts:
- Group SGLANG_CP_SHARED_KV_PER_LAYER_GROUP (default 8) consecutive layers into ONE
  RDMA submit: ~num_layers/K submits + events + enqueues instead of per-layer; same
  bytes (page index lists are identical across layers). on_layer_end is O(1) at
  non-boundary layers. The last partial group enqueues via the num_layers boundary;
  any misses fall back to one batched sync submit.
- Scheduler hook skips reqs already registered (bs>1 batch-forming re-iterates the
  same reqs ~9x -> was rebuilding the CP filter + context every time).

27 unit tests pass incl. grouping-boundary + batched-fallback.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-07 09:51:05 +00:00
co-authored by Claude Opus 4.8
parent 5fdf439c7e
commit e12afe8ced
4 changed files with 119 additions and 59 deletions
@@ -118,13 +118,15 @@ class TestPerLayerTransferContext(unittest.TestCase):
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.
# 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
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), 3) # layer 0 (overlap) + 1,2 (fallback)
self.assertEqual(eng.waits[-1], [101, 102, 103])
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
@@ -163,21 +165,24 @@ class TestBuildLayerBlocks(unittest.TestCase):
class _MockCtx:
num_layers = 8
def __init__(self):
self.submitted = []
self.submitted = [] # (start, end, event)
self.enqueued = 0
self.finished = False
self.failed = False
self._enq = set()
def note_enqueued(self, layer_id):
if layer_id in getattr(self, "_enq_layers", set()):
def note_enqueued(self, key):
if key in self._enq:
return False
self._enq_layers = getattr(self, "_enq_layers", set()) | {layer_id}
self._enq.add(key)
self.enqueued += 1
return True
def submit_layer(self, layer_id, event):
self.submitted.append((layer_id, event))
def submit_group(self, start, end, event):
self.submitted.append((start, end, event))
def finish(self):
self.finished = True
@@ -199,11 +204,14 @@ class _RecEvent:
self.synced = True
def _manager(num_workers=0):
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"
num_workers=num_workers,
event_factory=_RecEvent,
current_stream=lambda: "STREAM",
group_size=group_size,
)
@@ -221,23 +229,33 @@ def _drain(q):
class TestPerLayerTransferManager(unittest.TestCase):
def test_on_layer_end_enqueues_per_active_ctx_with_recorded_event(self):
m = _manager()
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, layer_id, ev in items:
self.assertEqual(layer_id, 5)
self.assertEqual(ev.recorded_stream, "STREAM") # event recorded on compute stream
for ctx, start, end, ev in items:
self.assertEqual((start, end), (5, 5))
self.assertEqual(ev.recorded_stream, "STREAM")
def test_worker_step_calls_submit_layer(self):
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_worker_step_calls_submit_group(self):
m = _manager()
c = _MockCtx()
ev = _RecEvent()
m._worker_step((c, 3, ev))
self.assertEqual(c.submitted, [(3, ev)])
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()
@@ -246,14 +264,14 @@ class TestPerLayerTransferManager(unittest.TestCase):
def __init__(self):
self.failed = False
def submit_layer(self, layer_id, event):
def submit_group(self, start, end, event):
raise RuntimeError("boom")
def mark_failed(self):
self.failed = True
b = _Boom()
m._worker_step((b, 0, None))
m._worker_step((b, 0, 0, None))
self.assertTrue(b.failed)
def test_finish_pops_and_calls_ctx_finish(self):
@@ -267,7 +285,7 @@ class TestPerLayerTransferManager(unittest.TestCase):
self.assertEqual(m.finish("r1"), 0) # idempotent after pop
def test_on_layer_end_no_active_is_noop(self):
m = _manager()
m = _manager(group_size=1)
m.on_layer_end(0)
self.assertEqual(_drain(m._q), [])