feat(disagg): per-layer transfer context for overlap (lever A, A2 core)

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>
This commit is contained in:
2026-06-07 00:31:58 +00:00
parent a09aec30a3
commit e0d47bfc41
2 changed files with 189 additions and 0 deletions

View File

@@ -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)

View File

@@ -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()