Add _transfer_layers_async: submit each layer's transfer non-blocking via the async API (batch_transfer_async_submit), pipelining layers in the RDMA engine, then wait for all once (wait_batch_transfers). Gated by the new SGLANG_CP_SHARED_KV_PER_LAYER_TRANSFER env (default off); replaces the monolithic all-layers batch_transfer_sync on that path. This is the transfer mechanism for per-layer overlap (lever A) and removes the per-layer blocking-sync tax measured in B1a; the forward-overlap hook (G2) builds on it next. Uses the safe async API, never the OnCuda busy-wait/_exit path. Unit-tested (test_per_layer_transfer.py, 5 cases): one submit per non-empty layer, single wait-for-all, empty-layer skip, submit-failure drain + return -1, and wait-status propagation. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
92 lines
3.2 KiB
Python
92 lines
3.2 KiB
Python
"""Unit tests for A1: the async per-layer transfer primitive
|
|
(MooncakeKVManager._transfer_layers_async).
|
|
|
|
It must submit each layer's transfer non-blocking (one batch_transfer_async_submit
|
|
per non-empty layer), then wait for all once; skip empty layers; and on a submit
|
|
failure, drain what was submitted and return -1. Engine is mocked (no RDMA/GPU).
|
|
"""
|
|
|
|
import unittest
|
|
|
|
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
|
|
|
|
|
class _FakeEngine:
|
|
def __init__(self, submit_rets=None, wait_ret=0):
|
|
self.submits = []
|
|
self.waits = []
|
|
self._submit_rets = submit_rets
|
|
self._wait_ret = wait_ret
|
|
self._n = 0
|
|
|
|
def batch_transfer_async_submit(self, sid, src, dst, lens):
|
|
self.submits.append((sid, list(src), list(dst), list(lens)))
|
|
if self._submit_rets is not None:
|
|
r = self._submit_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_ret
|
|
|
|
|
|
def _mgr(engine):
|
|
m = MooncakeKVManager.__new__(MooncakeKVManager) # skip heavy init
|
|
m.engine = engine
|
|
return m
|
|
|
|
|
|
def _one_block_per_layer(src, dst, item_len):
|
|
return [(src, dst, item_len)]
|
|
|
|
|
|
class TestPerLayerAsyncTransfer(unittest.TestCase):
|
|
def test_one_submit_per_layer_then_single_wait(self):
|
|
eng = _FakeEngine()
|
|
status = _mgr(eng)._transfer_layers_async(
|
|
"sess", [(1000, 2000, 64), (1064, 2064, 64), (1128, 2128, 64)],
|
|
_one_block_per_layer,
|
|
)
|
|
self.assertEqual(status, 0)
|
|
self.assertEqual(len(eng.submits), 3) # one async submit per layer
|
|
self.assertEqual(len(eng.waits), 1) # single wait-for-all
|
|
self.assertEqual(eng.waits[0], [101, 102, 103])
|
|
self.assertEqual(eng.submits[0], ("sess", [1000], [2000], [64]))
|
|
|
|
def test_skips_empty_layers(self):
|
|
eng = _FakeEngine()
|
|
|
|
def blocks(src, dst, il):
|
|
return [] if src == 1064 else [(src, dst, il)]
|
|
|
|
_mgr(eng)._transfer_layers_async(
|
|
"s", [(1000, 2000, 64), (1064, 2064, 64), (1128, 2128, 64)], blocks
|
|
)
|
|
self.assertEqual(len(eng.submits), 2) # middle layer skipped
|
|
self.assertEqual(eng.waits[0], [101, 102])
|
|
|
|
def test_submit_failure_drains_and_returns_minus_one(self):
|
|
eng = _FakeEngine(submit_rets=[101, -1]) # 2nd submit fails
|
|
status = _mgr(eng)._transfer_layers_async(
|
|
"s", [(1, 2, 3), (4, 5, 6), (7, 8, 9)], _one_block_per_layer
|
|
)
|
|
self.assertEqual(status, -1)
|
|
self.assertEqual(len(eng.submits), 2) # stopped at the failure
|
|
self.assertEqual(eng.waits, [[101]]) # drained the first submit
|
|
|
|
def test_wait_status_propagates(self):
|
|
eng = _FakeEngine(wait_ret=-7)
|
|
status = _mgr(eng)._transfer_layers_async("s", [(1, 2, 3)], _one_block_per_layer)
|
|
self.assertEqual(status, -7)
|
|
|
|
def test_no_layers_is_zero(self):
|
|
eng = _FakeEngine()
|
|
self.assertEqual(_mgr(eng)._transfer_layers_async("s", [], _one_block_per_layer), 0)
|
|
self.assertEqual(eng.waits, [[]])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|