Make CP HiCache backup admission deterministic

CP HiCache write-through under shared KV was still using rank-wide collectives to decide host reservation eviction, and per-layer backup registration could be bypassed before the final forward boundary. This moves backup registration to the final run_batch pre-forward boundary, forwards it through SessionAwareCache, exposes fallback paths as explicit warnings, and introduces deterministic owner-lane capacity planning for CP host reservation.

Constraint: CP shared-KV ranks must keep target and draft host reservations owner-lane consistent without adding hot-path collective synchronization

Constraint: Remote CUDA validation must run in the g0034 container, not locally

Rejected: Keep reserve_slots_max all_reduce as the default admission path | observed reserve collectives reaching double-digit and occasional 100ms+ latency

Rejected: Silent post-forward catch-up backup | hides when per-layer forward-overlap backup is not actually active

Confidence: medium

Scope-risk: broad

Directive: Do not reintroduce CP HiCache hot-path collectives without a measured mismatch case and explicit fallback warning

Tested: py_compile for modified Python modules and CP HiCache metadata test file in remote g0034 container

Tested: python3 -m pytest test/registered/unit/mem_cache/test_cp_hicache_metadata.py -q in remote g0034 container (75 passed, 5 warnings)

Tested: git diff --check HEAD~1..HEAD

Not-tested: Local pytest blocked by missing pybase64 in the local environment

Not-tested: Full CP HiCache + MTP E2E after the no-collective reservation change

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-05-27 23:11:23 +08:00
co-authored by OmX
parent f355fdd39e
commit 40a8de5fd1
11 changed files with 1476 additions and 118 deletions
@@ -108,6 +108,7 @@ from sglang.srt.mem_cache.hiradix_cache import (
_compute_shared_hicache_token_capacities,
)
from sglang.srt.mem_cache.radix_cache import RadixKey, TreeNode
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -597,6 +598,26 @@ class FakeEvictionStrategy:
return 0
class FakeCpLayout:
def __init__(self, cp_size=4, cp_rank=0, page_size=1):
self.cp_size = cp_size
self.cp_rank = cp_rank
self.page_size = page_size
def owner_for_logical_pages(self, logical_pages):
owners = torch.remainder(logical_pages - 1, self.cp_size)
return torch.where(logical_pages <= 0, torch.full_like(owners, -1), owners)
def logical_locs_to_physical(self, logical_locs):
return logical_locs
def owned_by_this_rank(self, logical_locs):
logical_pages = torch.div(
logical_locs, self.page_size, rounding_mode="floor"
)
return self.owner_for_logical_pages(logical_pages) == self.cp_rank
class FakeEvictDeviceController:
write_policy = "write_through"
@@ -610,6 +631,20 @@ class FakeTokenAllocator:
class TestHiRadixCacheCPBackup(CustomTestCase):
def test_session_aware_cache_forwards_cp_hicache_prepare(self):
calls = []
class Inner:
def prepare_write_backup_for_req(self, req):
calls.append(req)
wrapper = SessionAwareCache(Inner())
req = object()
wrapper.prepare_write_backup_for_req(req)
self.assertEqual(calls, [req])
def test_node_backuped_uses_cp_metadata(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
@@ -843,19 +878,21 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
self.assertEqual(node.hit_count, 1)
def test_write_backup_retries_by_required_physical_slots(self):
def test_write_backup_uses_deterministic_host_eviction_before_reserve(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.page_size = 1
reservation_factory = lambda device_indices: make_write_reservation(
device_indices, node_id=122
)
cache.cache_controller = FakeReserveWriteController(
[HiCacheWriteFailure(required_host_slots=1), reservation_factory]
)
cache.cache_controller = FakeReserveWriteController([reservation_factory])
cache.cache_controller.cp_shared_kv_layout = FakeCpLayout(cp_size=1, cp_rank=0)
cache.token_to_kv_pool_host = types.SimpleNamespace(size=16)
cache.evictable_host_leaves = set()
cache.eviction_strategy = FakeEvictionStrategy()
cache.get_child_key_fn = lambda key: key.token_ids[0]
cache._record_remove_event = lambda node: None
cache._update_host_leaf_status = lambda node: None
cache.ongoing_write_through = {}
cache.pending_host_backups = {}
cache.inc_node_lock_ref = lambda node: None
@@ -871,8 +908,8 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
evictable_node.host_len = 4
evictable_node.cp_hicache = CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([0], dtype=torch.int64),
host_indices=torch.tensor([55], dtype=torch.int64),
owned_positions=torch.arange(4, dtype=torch.int64),
host_indices=torch.arange(55, 59, dtype=torch.int64),
page_owners=torch.zeros(max(4, 0), dtype=torch.int8),
page_size=1,
)
@@ -886,7 +923,8 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
cache.write_backup(node)
self.assertEqual(
cache.cache_controller.evicted_host_indices[0].tolist(), [55]
cache.cache_controller.evicted_host_indices[0].tolist(),
list(range(55, 59)),
)
self.assertEqual(evictable_node.host_len, 0)
self.assertIsNone(evictable_node.cp_hicache)
@@ -896,51 +934,124 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
self.assertIn(node.id, cache.pending_host_backups)
self.assertEqual(len(cache.cache_controller.submitted), 1)
def test_write_backup_rolls_back_local_success_when_peer_needs_host_eviction(self):
def test_write_backup_cp_post_forward_path_warns_fallback(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.tp_world_size = 2
cache.tp_group = object()
cache.cache_controller = FakeReserveWriteController(
[
lambda device_indices: make_write_reservation(
device_indices, node_id=130, host_start=90
),
lambda device_indices: make_write_reservation(
device_indices, node_id=130, host_start=100
),
lambda device_indices, node_id: make_write_reservation(
device_indices, node_id=node_id
)
]
)
evictions = []
cache._evict_host_for_physical_slots = lambda required, synchronize_across_ranks=False: (
evictions.append((required, synchronize_across_ranks)) or required
)
cache.ongoing_write_through = {}
cache.pending_host_backups = {}
cache.inc_node_lock_ref = lambda node: None
reduce_values = iter([4, 0])
node = TreeNode()
node.id = 133
node.value = torch.arange(16, dtype=torch.int64)
def fake_all_reduce(tensor, op=None, group=None):
tensor.fill_(next(reduce_values))
with self.assertLogs(
"sglang.srt.mem_cache.hiradix_cache", level="WARNING"
) as captured:
cache.write_backup(node)
self.assertTrue(
any(
"[CP_HICACHE_FALLBACK][post_forward_catch_up_backup]" in message
for message in captured.output
)
)
def test_prepare_write_backup_for_req_chunked_skip_warns_fallback(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache.disable = False
cache._uses_cp_hicache = True
cache.cache_controller = FakeReserveWriteController([])
req = types.SimpleNamespace(
rid="rid-chunked",
is_chunked=1,
cp_hicache_prepared_backup=None,
)
with self.assertLogs(
"sglang.srt.mem_cache.hiradix_cache", level="WARNING"
) as captured:
cache.prepare_write_backup_for_req(req)
self.assertTrue(
any(
"[CP_HICACHE_FALLBACK][prepare_write_backup_skipped]" in message
and "reason=chunked_req" in message
for message in captured.output
)
)
def test_write_backup_deterministic_eviction_avoids_reserve_all_reduce(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.page_size = 1
cache.tp_world_size = 2
cache.tp_group = object()
cache.cache_controller = FakeReserveWriteController(
[
lambda device_indices: make_write_reservation(
device_indices, node_id=130, host_start=100
),
]
)
cache.cache_controller.cp_shared_kv_layout = FakeCpLayout(cp_size=1, cp_rank=0)
cache.token_to_kv_pool_host = types.SimpleNamespace(size=16)
evictions = []
cache._evict_host_for_physical_slots = lambda required, synchronize_across_ranks=False: (
evictions.append((required, synchronize_across_ranks)) or required
)
cache.evictable_host_leaves = set()
cache.eviction_strategy = FakeEvictionStrategy()
cache.get_child_key_fn = lambda key: key.token_ids[0]
cache._record_remove_event = lambda node: None
cache._update_host_leaf_status = lambda node: None
cache.ongoing_write_through = {}
cache.pending_host_backups = {}
cache.inc_node_lock_ref = lambda node: None
root = TreeNode()
root.key = RadixKey(token_ids=[], extra_key=None)
root.value = []
cache.root_node = root
evictable_node = TreeNode()
evictable_node.parent = root
evictable_node.key = RadixKey(token_ids=[2], extra_key=None)
evictable_node.value = None
evictable_node.host_len = 16
evictable_node.cp_hicache = CpHiCacheNodeMetadata(
logical_len=16,
owned_positions=torch.arange(16, dtype=torch.int64),
host_indices=torch.arange(90, 106, dtype=torch.int64),
page_owners=torch.zeros(16, dtype=torch.int8),
page_size=1,
)
root.children[2] = evictable_node
cache.evictable_host_leaves.add(evictable_node)
node = TreeNode()
node.id = 130
node.value = torch.arange(16, dtype=torch.int64)
with patch("torch.distributed.all_reduce", side_effect=fake_all_reduce):
with patch(
"torch.distributed.all_reduce",
side_effect=AssertionError("reserve admission must not all_reduce"),
):
backed_len = cache.write_backup(node)
self.assertEqual(backed_len, 16)
# The first local success must be released before all ranks enter the
# collective host eviction/retry branch. Otherwise CP ranks can diverge:
# successful ranks submit backup while failing ranks enter eviction
# all_reduce, which matches the observed Gloo 4-vs-1 mismatch.
self.assertEqual(
cache.cache_controller.evicted_host_indices[0].tolist(),
list(range(90, 106)),
)
self.assertEqual(evictions, [(4, True)])
self.assertEqual(evictions, [])
self.assertEqual(len(cache.cache_controller.submitted), 1)
self.assertEqual(
cache.cache_controller.submitted[0].metadata.host_indices.tolist(),
@@ -1111,7 +1222,7 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
self.assertEqual(cache.cache_controller.ack_write_queue, [])
self.assertEqual(evicted, [reservation.metadata])
def test_write_backup_cp_retry_failure_leaves_node_device_only(self):
def test_write_backup_cp_failfast_on_unplanned_reservation_failure(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = FakeReserveWriteController(
@@ -1130,9 +1241,12 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
node.id = 124
node.value = torch.arange(16, dtype=torch.int64)
backed_len = cache.write_backup(node)
with self.assertRaisesRegex(
RuntimeError,
"planner predicted no eviction need",
):
cache.write_backup(node)
self.assertEqual(backed_len, 0)
self.assertEqual(node.host_len, 0)
self.assertIsNone(node.cp_hicache)
self.assertEqual(cache.pending_host_backups, {})
@@ -1229,6 +1343,109 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
self.assertIsNotNone(node.cp_hicache)
self.assertEqual(node.cp_hicache.host_indices.tolist(), [55])
def test_cp_owner_token_counts_use_page_owners_and_page_size(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=4, cp_rank=0)
)
counts = cache._cp_owner_token_counts(
torch.tensor([0, 1, 1, 2], dtype=torch.int8),
page_size=64,
)
self.assertEqual(counts, (64, 128, 64, 0))
def test_cp_metadata_count_asserts_local_owner_lane(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=3, cp_rank=1),
has_draft_hicache=True,
)
metadata = CpHiCacheNodeMetadata(
logical_len=192,
owned_positions=torch.arange(64, dtype=torch.int64),
host_indices=torch.arange(100, 164, dtype=torch.int64),
draft_host_indices=torch.arange(200, 264, dtype=torch.int64),
page_owners=torch.tensor([0, 1, 2], dtype=torch.int8),
page_size=64,
)
counts = cache._cp_assert_metadata_counts(metadata, context="unit")
self.assertEqual(counts, (64, 64, 64))
def test_cp_metadata_count_asserts_mismatched_local_slots(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=2, cp_rank=1),
has_draft_hicache=False,
)
metadata = CpHiCacheNodeMetadata(
logical_len=128,
owned_positions=torch.empty((0,), dtype=torch.int64),
host_indices=torch.empty((0,), dtype=torch.int64),
page_owners=torch.tensor([0, 1], dtype=torch.int8),
page_size=64,
)
with self.assertRaisesRegex(RuntimeError, "owner lane mismatch"):
cache._cp_assert_metadata_counts(metadata, context="unit")
def test_cp_host_capacity_snapshot_counts_committed_and_pending_draft(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=2, cp_rank=0),
has_draft_hicache=True,
draft_mem_pool_host=types.SimpleNamespace(size=512),
)
cache.token_to_kv_pool_host = types.SimpleNamespace(size=512)
cache.pending_host_backups = {}
root = TreeNode()
root.children = {}
cache.root_node = root
committed = TreeNode()
committed.id = 501
committed.parent = root
committed.key = RadixKey([1])
committed.host_len = 128
committed.cp_hicache = CpHiCacheNodeMetadata(
logical_len=128,
owned_positions=torch.arange(64, dtype=torch.int64),
host_indices=torch.arange(64, dtype=torch.int64),
draft_host_indices=torch.arange(100, 164, dtype=torch.int64),
page_owners=torch.tensor([0, 1], dtype=torch.int8),
page_size=64,
)
root.children[1] = committed
pending = TreeNode()
pending.id = 502
pending_metadata = CpHiCacheNodeMetadata(
logical_len=64,
owned_positions=torch.arange(64, dtype=torch.int64),
host_indices=torch.arange(200, 264, dtype=torch.int64),
draft_host_indices=torch.arange(300, 364, dtype=torch.int64),
page_owners=torch.tensor([0], dtype=torch.int8),
page_size=64,
)
cache.pending_host_backups[pending.id] = PendingHiCacheBackup(
node=pending,
metadata=pending_metadata,
logical_len=64,
)
snapshot = cache._cp_host_capacity_snapshot()
self.assertEqual(snapshot.target_capacity, (512, 512))
self.assertEqual(snapshot.draft_capacity, (512, 512))
self.assertEqual(snapshot.committed_target, (64, 64))
self.assertEqual(snapshot.committed_draft, (64, 64))
self.assertEqual(snapshot.pending_target, (64, 0))
self.assertEqual(snapshot.pending_draft, (64, 0))
class TestHiRadixCacheCPSplitEvict(CustomTestCase):
def test_split_node_splits_cp_metadata_by_owned_positions(self):
@@ -1558,6 +1775,92 @@ class TestHiRadixCacheCPSplitEvict(CustomTestCase):
self.assertEqual(node.host_len, 4)
self.assertIsNotNone(node.cp_hicache)
def test_cp_host_eviction_plan_is_stable_across_leaf_insertion_order(self):
def make_cache(order):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=2, cp_rank=0),
has_draft_hicache=False,
)
cache.eviction_strategy = FakeEvictionStrategy()
cache.pending_host_backups = {}
cache.root_node = TreeNode()
cache.root_node.children = {}
cache.get_child_key_fn = lambda key: key.token_ids[0]
cache._clear_pin = lambda node: None
nodes = []
for node_id in (42, 41):
node = TreeNode(id=node_id)
node.parent = cache.root_node
node.key = RadixKey([node_id])
node.value = None
node.lock_ref = 0
node.host_ref_counter = 0
node.host_len = 64
node.cp_hicache = CpHiCacheNodeMetadata(
logical_len=64,
owned_positions=torch.arange(64, dtype=torch.int64),
host_indices=torch.arange(node_id * 100, node_id * 100 + 64),
page_owners=torch.tensor([0], dtype=torch.int8),
page_size=64,
)
cache.root_node.children[node_id] = node
nodes.append(node)
cache.evictable_host_leaves = {nodes[i] for i in order}
return cache
first = make_cache([0, 1])._plan_cp_host_evictions((64, 0))
second = make_cache([1, 0])._plan_cp_host_evictions((64, 0))
self.assertEqual([node.id for node in first.victims], [41])
self.assertEqual([node.id for node in second.victims], [41])
self.assertEqual(first.planned_freed, (64, 0))
self.assertEqual(first.remaining_deficit, (0, 0))
def test_cp_host_eviction_plan_skips_pending_backup_nodes(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = types.SimpleNamespace(
cp_shared_kv_layout=FakeCpLayout(cp_size=1, cp_rank=0),
has_draft_hicache=False,
)
cache.eviction_strategy = FakeEvictionStrategy()
cache.root_node = TreeNode()
cache.root_node.children = {}
cache.get_child_key_fn = lambda key: key.token_ids[0]
cache.pending_host_backups = {}
cache.evictable_host_leaves = set()
for node_id in (10, 11):
node = TreeNode(id=node_id)
node.parent = cache.root_node
node.key = RadixKey([node_id])
node.value = None
node.lock_ref = 0
node.host_ref_counter = 0
node.host_len = 64
node.cp_hicache = CpHiCacheNodeMetadata(
logical_len=64,
owned_positions=torch.arange(64, dtype=torch.int64),
host_indices=torch.arange(node_id * 100, node_id * 100 + 64),
page_owners=torch.tensor([0], dtype=torch.int8),
page_size=64,
)
cache.root_node.children[node_id] = node
cache.evictable_host_leaves.add(node)
if node_id == 10:
cache.pending_host_backups[node.id] = PendingHiCacheBackup(
node=node,
metadata=node.cp_hicache,
logical_len=64,
)
plan = cache._plan_cp_host_evictions((64,))
self.assertEqual([node.id for node in plan.victims], [11])
self.assertEqual(plan.planned_freed, (64,))
class TestHiRadixCacheCPLoadBack(CustomTestCase):
def test_cp_load_back_uses_host_len_not_host_value(self):