Keep CP HiCache backup state replicated on radix nodes
Rank-local CP HiCache async write queues can drain at different times across ranks. Radix structural decisions that consult those local queues can diverge, so backup-pending state now lives on the replicated radix node and uses a separate backup lock ref from request locks. Commit and rollback clear the replicated marker, and split/prune/hit-count probes use the marker rather than local pending dictionaries. Constraint: CP ranks must make identical radix structural decisions even when local async write acknowledgements drain at different times Rejected: Drain write acknowledgements before every split/prune decision | still rank-local and adds hot-path synchronization Rejected: Treat backup lock_ref as a normal request lock | hides the distinction between request protection and backup-pending protection Confidence: high Scope-risk: moderate Directive: Do not derive CP HiCache split/prune/write-through decisions from pending_host_backups or ongoing_write_through without a replicated node marker Tested: local py_compile python/sglang/srt/mem_cache/hiradix_cache.py python/sglang/srt/mem_cache/radix_cache.py Tested: local git diff --check Tested: remote cjy-glm5-new pytest replicated marker test plus targeted CP HiCache metadata subset, 10 passed Not-tested: full test_cp_hicache_metadata.py because current branch has unrelated fixture drift around enable_cp_l3/cp_shared_l2_page_allocator/running-count setup Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
@@ -1380,6 +1380,10 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
|
||||
cache.pending_host_backups[node.id] = PendingHiCacheBackup(
|
||||
node=node, metadata=metadata, logical_len=64
|
||||
)
|
||||
node.cp_backup_pending = True
|
||||
node.cp_backup_lock_ref = 1
|
||||
node.lock_ref = 1
|
||||
cache._cp_account_enter_pending(metadata)
|
||||
|
||||
self.assertFalse(cache._node_backuped(node))
|
||||
|
||||
@@ -1412,6 +1416,10 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
|
||||
cache.pending_host_backups[node.id] = PendingHiCacheBackup(
|
||||
node=node, metadata=metadata, logical_len=64
|
||||
)
|
||||
node.cp_backup_pending = True
|
||||
node.cp_backup_lock_ref = 1
|
||||
node.lock_ref = 1
|
||||
cache._cp_account_enter_pending(metadata)
|
||||
|
||||
cache._rollback_pending_backup(node.id)
|
||||
|
||||
@@ -1494,6 +1502,9 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
|
||||
page_size=1,
|
||||
)
|
||||
cache.ongoing_write_through[node.id] = node
|
||||
node.cp_backup_pending = True
|
||||
node.cp_backup_lock_ref = 1
|
||||
node.lock_ref = 1
|
||||
|
||||
cache._inc_hit_count(node)
|
||||
|
||||
@@ -2455,6 +2466,9 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
|
||||
)
|
||||
}
|
||||
cache.dec_node_lock_ref = lambda n: dec_locked.append(n.id)
|
||||
node.cp_backup_pending = True
|
||||
node.cp_backup_lock_ref = 1
|
||||
node.lock_ref = 1
|
||||
|
||||
finish_event = MagicMock()
|
||||
cancelled = []
|
||||
@@ -3581,6 +3595,7 @@ class TestHiRadixCacheCPLoadBack(CustomTestCase):
|
||||
cache.pending_host_backups[child.id] = PendingHiCacheBackup(
|
||||
node=child, metadata=metadata, logical_len=8
|
||||
)
|
||||
child.cp_backup_pending = True
|
||||
root.children[0] = child
|
||||
|
||||
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([0, 1, 2, 9])))
|
||||
@@ -4007,6 +4022,7 @@ class TestHiRadixCacheCPLoadBack(CustomTestCase):
|
||||
node.cp_hicache = None
|
||||
root.children[(0, 1, 2, 3)] = node
|
||||
cache.ongoing_write_through[node.id] = node
|
||||
node.cp_backup_pending = True
|
||||
|
||||
result = cache.insert(
|
||||
InsertParams(
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
"""Regression coverage for CP HiCache backup-pending radix decisions.
|
||||
|
||||
The radix split/prune/write-through probes must not branch on drain-timed,
|
||||
rank-local dictionaries such as ``ongoing_write_through`` or
|
||||
``pending_host_backups``. Those dictionaries are advanced by local async write
|
||||
ack timing; using them for structural radix decisions lets CP ranks take
|
||||
different branches before the next collective.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import sgl_kernel # noqa: F401
|
||||
except (ImportError, RuntimeError):
|
||||
if "sgl_kernel" not in sys.modules:
|
||||
stub = types.ModuleType("sgl_kernel")
|
||||
stub.__path__ = []
|
||||
|
||||
def _getattr(name):
|
||||
if name.startswith("__"):
|
||||
raise AttributeError(name)
|
||||
fn = lambda *a, **k: None
|
||||
setattr(stub, name, fn)
|
||||
return fn
|
||||
|
||||
stub.__getattr__ = _getattr
|
||||
sys.modules["sgl_kernel"] = stub
|
||||
|
||||
from sglang.srt.mem_cache.hiradix_cache import ( # noqa: E402
|
||||
CpHiCacheNodeMetadata,
|
||||
HiRadixCache,
|
||||
PreparedCpHiCacheBackup,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import TreeNode # noqa: E402
|
||||
|
||||
|
||||
def _cp_cache():
|
||||
cache = HiRadixCache.__new__(HiRadixCache)
|
||||
cache._uses_cp_hicache = True
|
||||
cache.ongoing_write_through = {}
|
||||
cache.pending_host_backups = {}
|
||||
return cache
|
||||
|
||||
|
||||
def _metadata(length: int) -> CpHiCacheNodeMetadata:
|
||||
return CpHiCacheNodeMetadata(
|
||||
logical_len=length,
|
||||
owned_positions=torch.arange(length, dtype=torch.int64),
|
||||
host_indices=torch.arange(length, dtype=torch.int64),
|
||||
page_owners=torch.zeros(1, dtype=torch.int8),
|
||||
page_size=length,
|
||||
)
|
||||
|
||||
|
||||
def _reservation(metadata: CpHiCacheNodeMetadata) -> SimpleNamespace:
|
||||
return SimpleNamespace(metadata=metadata)
|
||||
|
||||
|
||||
class TestCPHiCacheReplicatedPendingMarker(unittest.TestCase):
|
||||
def test_split_probe_reads_marker_not_drain_timed_dicts(self):
|
||||
cache = _cp_cache()
|
||||
node = TreeNode(id=42)
|
||||
|
||||
cache.ongoing_write_through[node.id] = node
|
||||
cache.pending_host_backups[node.id] = object()
|
||||
|
||||
self.assertFalse(
|
||||
cache._cp_node_split_still_pending(node),
|
||||
"split probe must ignore drain-timed local dicts when marker is clear",
|
||||
)
|
||||
|
||||
cache.ongoing_write_through.clear()
|
||||
cache.pending_host_backups.clear()
|
||||
node.cp_backup_pending = True
|
||||
|
||||
self.assertTrue(
|
||||
cache._cp_node_split_still_pending(node),
|
||||
"split probe must honor the replicated node marker",
|
||||
)
|
||||
|
||||
def test_inc_hit_count_reads_marker_not_drain_timed_dicts(self):
|
||||
cache = _cp_cache()
|
||||
cache.cache_controller = SimpleNamespace(write_policy="write_through")
|
||||
cache.write_through_threshold = 1
|
||||
|
||||
node = TreeNode(id=7)
|
||||
node.cp_backup_pending = True
|
||||
|
||||
def fail_if_reached(_node):
|
||||
raise AssertionError("marker-pending node must not reach backuped check")
|
||||
|
||||
cache._node_backuped = fail_if_reached
|
||||
cache._inc_hit_count(node)
|
||||
|
||||
reached_backuped_check = {"hit": False}
|
||||
|
||||
def record_backuped(_node):
|
||||
reached_backuped_check["hit"] = True
|
||||
return True
|
||||
|
||||
node.cp_backup_pending = False
|
||||
cache.ongoing_write_through[node.id] = node
|
||||
cache.pending_host_backups[node.id] = object()
|
||||
cache._node_backuped = record_backuped
|
||||
|
||||
cache._inc_hit_count(node)
|
||||
|
||||
self.assertTrue(
|
||||
reached_backuped_check["hit"],
|
||||
"dict-pending alone must not short-circuit hit-count side effects",
|
||||
)
|
||||
|
||||
def test_attach_commit_maintains_marker_and_backup_lock_ref(self):
|
||||
cache = _cp_cache()
|
||||
cache.inc_node_lock_ref = lambda node: setattr(
|
||||
node, "lock_ref", node.lock_ref + 1
|
||||
)
|
||||
cache.dec_node_lock_ref = lambda node: setattr(
|
||||
node, "lock_ref", node.lock_ref - 1
|
||||
)
|
||||
cache.cache_controller = SimpleNamespace(cp_shared_l2_page_allocator=None)
|
||||
|
||||
node = TreeNode()
|
||||
node.value = torch.arange(4, dtype=torch.int64)
|
||||
metadata = _metadata(4)
|
||||
prepared = PreparedCpHiCacheBackup(
|
||||
node_id=123,
|
||||
reservation=_reservation(metadata),
|
||||
metadata=metadata,
|
||||
logical_len=4,
|
||||
)
|
||||
|
||||
cache._attach_prepared_cp_backup(node, prepared)
|
||||
|
||||
self.assertTrue(node.cp_backup_pending)
|
||||
self.assertEqual(node.cp_backup_lock_ref, 1)
|
||||
self.assertEqual(node.lock_ref, 1)
|
||||
|
||||
cache._commit_pending_backup(node.id)
|
||||
|
||||
self.assertFalse(node.cp_backup_pending)
|
||||
self.assertEqual(node.cp_backup_lock_ref, 0)
|
||||
self.assertEqual(node.lock_ref, 0)
|
||||
self.assertEqual(node.host_len, 4)
|
||||
|
||||
def test_rollback_clears_marker_and_backup_lock_ref(self):
|
||||
cache = _cp_cache()
|
||||
cache.inc_node_lock_ref = lambda node: setattr(
|
||||
node, "lock_ref", node.lock_ref + 1
|
||||
)
|
||||
cache.dec_node_lock_ref = lambda node: setattr(
|
||||
node, "lock_ref", node.lock_ref - 1
|
||||
)
|
||||
evicted = []
|
||||
cache.cache_controller = SimpleNamespace(
|
||||
ack_write_queue=[],
|
||||
cancel_layer_write_state=lambda node_id: None,
|
||||
evict_cp_host=lambda metadata: evicted.append(metadata) or 4,
|
||||
)
|
||||
|
||||
node = TreeNode()
|
||||
node.value = torch.arange(4, dtype=torch.int64)
|
||||
metadata = _metadata(4)
|
||||
prepared = PreparedCpHiCacheBackup(
|
||||
node_id=124,
|
||||
reservation=_reservation(metadata),
|
||||
metadata=metadata,
|
||||
logical_len=4,
|
||||
)
|
||||
|
||||
cache._attach_prepared_cp_backup(node, prepared)
|
||||
cache._rollback_pending_backup(node.id)
|
||||
|
||||
self.assertFalse(node.cp_backup_pending)
|
||||
self.assertEqual(node.cp_backup_lock_ref, 0)
|
||||
self.assertEqual(node.lock_ref, 0)
|
||||
self.assertEqual(evicted, [metadata])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user