refactor(hicache): flatten L2 transfer execution (#34793)

GB300 test fails unrelated
This commit is contained in:
cctry
2026-08-16 00:33:34 -07:00
committed by GitHub
parent 56a759cffc
commit 8922bb98e2
10 changed files with 767 additions and 648 deletions
@@ -16,12 +16,14 @@ register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small")
class TestLoadBackDurationMetric(CustomTestCase):
def setUp(self):
from sglang.srt.managers import cache_controller as cc
from sglang.srt.mem_cache import l2_transfer as transfer
cc._timing_events_supported.cache_clear()
transfer._timing_events_supported.cache_clear()
self.cc = cc
self.transfer = transfer
def _completed_pair(self, payload_floats=1024 * 1024):
start, finish, timing_enabled = self.cc.make_timing_event_pair()
start, finish, timing_enabled = self.transfer.make_timing_event_pair()
self.assertTrue(timing_enabled)
stream = torch.cuda.Stream()
start.record()
@@ -46,9 +48,11 @@ class TestLoadBackDurationMetric(CustomTestCase):
events.append(event)
return event
with patch.object(self.cc.device_module, "Event", side_effect=create_event):
self.cc._timing_events_supported.cache_clear()
start, finish, timing_enabled = self.cc.make_timing_event_pair()
with patch.object(
self.transfer.device_module, "Event", side_effect=create_event
):
self.transfer._timing_events_supported.cache_clear()
start, finish, timing_enabled = self.transfer.make_timing_event_pair()
self.assertFalse(timing_enabled)
self.assertIs(start, events[0])
@@ -7,17 +7,17 @@ from unittest import mock
import torch
from sglang.srt.managers import cache_controller as manager_cache_controller
from sglang.srt.managers.cache_controller import CacheOperation as ManagerCacheOperation
from sglang.srt.managers.cache_controller import (
HiCacheController,
from sglang.srt.managers.cache_controller import CacheOperation, HiCacheController
from sglang.srt.mem_cache import l2_transfer as transfer_module
from sglang.srt.mem_cache.hicache_storage import (
PoolHitPolicy,
PoolName,
PoolTransfer,
)
from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer
from sglang.srt.mem_cache.hybrid_cache import hybrid_cache_controller
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
CacheOperation,
HybridCacheController,
)
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
from sglang.srt.mem_cache.memory_pool_host import (
DeepSeekV4PagedHostPool,
DeepSeekV4StateHostPool,
@@ -30,6 +30,7 @@ from sglang.srt.mem_cache.memory_pool_host import (
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
@@ -59,6 +60,33 @@ def _device_pool_stub(*, layer_num: int, **fields) -> SimpleNamespace:
)
def _host_group_stub(captured, *, can_use_write_back_jit: bool) -> SimpleNamespace:
class FakeHostPool:
size_per_token = 2
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
captured.append(host_indices)
entries = [
PoolEntry(
name=name,
host_pool=FakeHostPool(),
device_pool=None,
layer_mapper=lambda layer_id: layer_id,
is_primary_index_anchor=name == PoolName.KV,
)
for name in (PoolName.KV, PoolName.SWA, PoolName.DEEPSEEK_V4_C4)
]
return SimpleNamespace(
layout="page_first",
can_use_write_back_jit=can_use_write_back_jit,
anchor_entry=entries[0],
entry_map={entry.name: entry for entry in entries},
)
def _cpu_staged_lf_pf_copy(
src_registry,
*,
@@ -152,21 +180,204 @@ class _FakeEvent:
class _FakeDeviceModule:
Event = _FakeEvent
@staticmethod
def Stream():
return object()
@staticmethod
@contextmanager
def stream(stream):
yield
class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
def setUp(self):
# start_writing probes timing support via a module-cached check;
# clear it on both sides so results from (or against) the fake
# device module never leak across tests.
manager_cache_controller._timing_events_supported.cache_clear()
transfer_module._timing_events_supported.cache_clear()
self.addCleanup(transfer_module._timing_events_supported.cache_clear)
def tearDown(self):
manager_cache_controller._timing_events_supported.cache_clear()
@staticmethod
def _start_writing(controller):
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
controller.l2_transfer_engine = L2TransferEngine("kernel")
controller.start_writing()
def test_hybrid_load_forwards_merged_pool_transfers(self):
transfer = PoolTransfer(
name=PoolName.SWA,
host_indices=_indices(0, 2),
device_indices=_indices(2, 4),
keys=["page-key"],
hit_policy=PoolHitPolicy.TRAILING_PAGES,
)
op = CacheOperation(_indices(0, 4), _indices(4, 8), 7)
op.pool_transfers = [transfer]
controller = mock.Mock(spec=HybridCacheController)
controller.load_queue = [op, op]
controller.layer_done_counter = mock.MagicMock()
controller.layer_done_counter.update_producer.return_value = 0
controller._move_op_indices.side_effect = lambda op: (
op.host_indices,
op.device_indices,
op.pool_transfers,
)
controller.mem_pool_host = _host_group_stub([], can_use_write_back_jit=False)
controller.has_draft = False
controller.has_mtp_draft = False
controller._l2_transfers.side_effect = lambda *args: (
HybridCacheController._l2_transfers(controller, *args)
)
controller._l2_load_transfers.side_effect = lambda *args: (
HybridCacheController._l2_load_transfers(controller, *args)
)
controller._num_tokens_by_pool.return_value = {}
controller._transfer_num_bytes.return_value = 0
controller.l2_transfer_engine = mock.Mock()
completion = SimpleNamespace(
start_event=object(), finish_event=object(), timing_enabled=False
)
controller.l2_transfer_engine.submit_host_to_device.return_value = completion
controller.layer_num = 2
controller.ack_load_queue = []
self.assertEqual(HybridCacheController.start_loading(controller), 0)
merged_op = controller._move_op_indices.call_args.args[0]
merged_transfer = merged_op.pool_transfers[0]
self.assertEqual(merged_transfer.host_indices.tolist(), [0, 1, 0, 1])
self.assertEqual(merged_transfer.keys, ["page-key", "page-key"])
self.assertEqual(merged_transfer.hit_policy, PoolHitPolicy.TRAILING_PAGES)
controller._l2_load_transfers.assert_called_once()
l2_transfers = (
controller.l2_transfer_engine.submit_host_to_device.call_args.args[0]
)
self.assertEqual(len(l2_transfers), 2)
self.assertEqual(l2_transfers[1].host_indices.tolist(), [0, 1, 0, 1])
self.assertEqual(
len(
HybridCacheController._l2_transfers(
controller, _indices(0, 0), _indices(0, 0), [merged_transfer]
)
),
1,
)
controller._num_tokens_by_pool.assert_called_once_with(merged_op)
self.assertEqual(controller.ack_load_queue[0].node_ids, [7, 7])
def test_l2_transfer_maps_global_layers(self):
host_pool = mock.Mock()
transfer = L2Transfer(
host_pool=host_pool,
device_pool=mock.sentinel.device_pool,
host_indices=_indices(0, 2),
device_indices=_indices(2, 4),
layer_mapper={1: 0, 3: 1}.get,
)
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
L2TransferEngine("kernel").submit_host_to_device([transfer], layer_num=4)
self.assertEqual(
[
call.args[3]
for call in host_pool.load_to_device_per_layer.call_args_list
],
[0, 1],
)
def test_packed_draft_load_is_flattened_into_l2_transfers(self):
host_pool = mock.Mock()
controller = HybridCacheController.__new__(HybridCacheController)
controller.mem_pool_host = SimpleNamespace(
anchor_entry=PoolEntry(
name=PoolName.KV,
host_pool=host_pool,
device_pool=mock.sentinel.target_device_pool,
layer_mapper={0: 0, 1: 1, 2: 2}.get,
is_primary_index_anchor=True,
),
entry_map={},
)
controller.layer_num = 2
controller.has_mtp_draft = True
controller.mtp_draft_device_pools = (mock.sentinel.draft_device_pool,)
controller.has_draft = False
self.assertEqual(
len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1
)
transfers = controller._l2_load_transfers(_indices(0, 2), _indices(2, 4))
self.assertEqual(len(transfers), 2)
self.assertFalse(transfers[0].is_draft)
self.assertTrue(transfers[1].is_draft)
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
L2TransferEngine("kernel").submit_host_to_device(transfers, layer_num=2)
self.assertEqual(
[
call.args[3]
for call in host_pool.load_to_device_per_layer.call_args_list
],
[0, 2, 1],
)
self.assertIs(
host_pool.load_to_device_per_layer.call_args_list[1].args[0],
mock.sentinel.draft_device_pool,
)
self.assertTrue(
host_pool.load_to_device_per_layer.call_args_list[1].kwargs["is_draft"]
)
def test_mixed_staged_write_resolves_indices_per_pool(self):
anchor_host_pool = SimpleNamespace(can_use_write_back_jit=True)
extra_host_pool = SimpleNamespace(can_use_write_back_jit=False)
anchor_entry = PoolEntry(
name=PoolName.KV,
host_pool=anchor_host_pool,
device_pool=None,
layer_mapper=lambda layer_id: layer_id,
is_primary_index_anchor=True,
)
extra_entry = PoolEntry(
name=PoolName.SWA,
host_pool=extra_host_pool,
device_pool=None,
layer_mapper=lambda layer_id: layer_id,
)
host_group = SimpleNamespace(
layout="page_first",
can_use_write_back_jit=False,
supports_per_pool_backup_indices=True,
anchor_entry=anchor_entry,
entry_map={PoolName.KV: anchor_entry, PoolName.SWA: extra_entry},
)
transfer = PoolTransfer(
name=PoolName.SWA,
host_indices=_indices(4, 6),
device_indices=_indices(6, 8),
)
op = CacheOperation(
host_indices=_indices(0, 2),
device_indices=_indices(2, 4),
node_id=1,
pool_transfers=[transfer],
)
controller = HybridCacheController.__new__(HybridCacheController)
controller.io_backend = "kernel"
controller.mem_pool_host = host_group
controller.move_indices = mock.Mock(
return_value=(mock.sentinel.host_indices, mock.sentinel.device_indices)
)
host_indices, device_indices, pool_transfers = controller._move_write_operation(
op
)
self.assertIs(host_indices, op.host_indices)
self.assertIs(device_indices, op.device_indices)
controller.move_indices.assert_called_once_with(
transfer.host_indices, transfer.device_indices
)
self.assertIs(pool_transfers[0].host_indices, mock.sentinel.host_indices)
self.assertIs(pool_transfers[0].device_indices, mock.sentinel.device_indices)
def _patched_transfers(self, src_registry=None, module=MEMORY_POOL_HOST_MODULE):
staged_side_effect = None
@@ -701,26 +912,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
self.assertIsNone(group.destroy())
def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self):
captured = {}
class FakeHostGroup:
layout = "page_first"
can_use_write_back_jit = True
anchor_entry = SimpleNamespace(
name=PoolName.KV, host_pool=SimpleNamespace(size_per_token=2)
)
entry_map = {}
def backup_from_device_all_layer(
self,
device_pool,
host_indices,
device_indices,
io_backend,
pool_transfers=None,
):
captured["host_indices"] = host_indices
captured["pool_transfers"] = pool_transfers
captured = []
controller = HybridCacheController.__new__(HybridCacheController)
controller.write_queue = [
@@ -738,53 +930,25 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
)
]
controller.io_backend = "kernel"
controller.mem_pool_host = FakeHostGroup()
controller.mem_pool_host = _host_group_stub(
captured, can_use_write_back_jit=True
)
controller.mem_pool_device = None
controller.has_draft = False
controller.write_stream = object()
controller.ack_write_queue = []
controller._record_transfer_indices_on_stream = lambda *args: None
controller.move_hybrid_indices = mock.Mock(
side_effect=AssertionError(
"write-back JIT kernel write should not move indices"
)
)
with (
mock.patch.object(
hybrid_cache_controller, "device_module", _FakeDeviceModule
),
mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
),
):
controller.start_writing()
self._start_writing(controller)
controller.move_hybrid_indices.assert_not_called()
self.assertEqual(captured["host_indices"].device.type, "cpu")
self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu")
self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"])
def test_hybrid_write_moves_indices_without_write_back_jit(self):
captured = {}
class FakeHostGroup:
layout = "page_first"
can_use_write_back_jit = False
anchor_entry = SimpleNamespace(
name=PoolName.KV, host_pool=SimpleNamespace(size_per_token=2)
)
entry_map = {}
def backup_from_device_all_layer(
self,
device_pool,
host_indices,
device_indices,
io_backend,
pool_transfers=None,
):
captured["host_indices"] = host_indices
captured["pool_transfers"] = pool_transfers
captured = []
op = CacheOperation(
host_indices=_indices(0, 4),
@@ -801,29 +965,20 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller = HybridCacheController.__new__(HybridCacheController)
controller.write_queue = [op]
controller.io_backend = "kernel"
controller.mem_pool_host = FakeHostGroup()
controller.mem_pool_host = _host_group_stub(
captured, can_use_write_back_jit=False
)
controller.mem_pool_device = None
controller.has_draft = False
controller.write_stream = object()
controller.ack_write_queue = []
controller._record_transfer_indices_on_stream = lambda *args: None
controller.move_hybrid_indices = mock.Mock(
return_value=(op.host_indices, op.device_indices, op.pool_transfers)
)
with (
mock.patch.object(
hybrid_cache_controller, "device_module", _FakeDeviceModule
),
mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
),
):
controller.start_writing()
self._start_writing(controller)
controller.move_hybrid_indices.assert_called_once()
self.assertEqual(captured["host_indices"].device.type, "cpu")
self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu")
self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"])
def test_write_back_jit_cache_controller_keeps_host_indices_on_cpu(self):
captured = {}
@@ -840,7 +995,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller = HiCacheController.__new__(HiCacheController)
controller.write_queue = [
ManagerCacheOperation(
CacheOperation(
host_indices=_indices(0, 4),
device_indices=_indices(4, 8),
node_id=1,
@@ -850,7 +1005,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller.mem_pool_host = FakeHostPool()
controller.mem_pool_device = None
controller.has_draft = False
controller.write_stream = object()
controller.device = "cuda"
controller.ack_write_queue = []
controller.move_indices = mock.Mock(
side_effect=AssertionError(
@@ -858,10 +1013,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
)
)
with mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
):
controller.start_writing()
self._start_writing(controller)
controller.move_indices.assert_not_called()
self.assertEqual(captured["host_indices"].device.type, "cpu")
@@ -879,7 +1031,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
):
captured["host_indices"] = host_indices
op = ManagerCacheOperation(
op = CacheOperation(
host_indices=_indices(0, 4),
device_indices=_indices(4, 8),
node_id=1,
@@ -890,16 +1042,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller.mem_pool_host = FakeHostPool()
controller.mem_pool_device = None
controller.has_draft = False
controller.write_stream = object()
controller.device = "cuda"
controller.ack_write_queue = []
controller.move_indices = mock.Mock(
return_value=(op.host_indices, op.device_indices)
)
with mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
):
controller.start_writing()
self._start_writing(controller)
controller.move_indices.assert_called_once()
self.assertEqual(captured["host_indices"].device.type, "cpu")
@@ -148,7 +148,6 @@ _EXPOSED = {
("disaggregation/common/conn.py", "disaggregation_bootstrap_port"),
("disaggregation/common/conn.py", "pp_size"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_mem_layout"),
("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"),
("disaggregation/encode_receiver.py", "disaggregation_ib_device"),
("disaggregation/encode_receiver.py", "encoder_transfer_backend"),