fix: synchronize CP HiCache host eviction semantics

This commit is contained in:
2026-05-08 03:20:47 +08:00
parent eecc8e21ec
commit 8844c303e2
4 changed files with 144 additions and 6 deletions
@@ -246,6 +246,13 @@ class TestHiCacheControllerCPWrite(CustomTestCase):
self.assertEqual(config.model_name, "test-model")
self.assertEqual(config.tp_lcm_size, 8)
def test_attach_storage_backend_rejects_cp_hicache(self):
host_pool = FakeHostPool(torch.tensor([], dtype=torch.int64))
controller = self.make_controller(host_pool)
with self.assertRaisesRegex(RuntimeError, "CP shared KV.*storage backend"):
controller.attach_storage_backend("mooncake")
class TestHiCacheControllerCPLoad(TestHiCacheControllerCPWrite):
def test_cp_load_allocates_full_logical_locs_and_transfers_owned_physical_locs(self):
@@ -1,6 +1,6 @@
import sys
import unittest
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
import torch
@@ -197,6 +197,19 @@ class FakeWriteController:
return len(host_indices)
class FakeZeroOwnedWriteController:
write_policy = "write_through"
def write(self, device_indices, node_id=-1, priority=None):
return FakeWriteSuccess(
CpHiCacheNodeMetadata(
logical_len=len(device_indices),
owned_positions=torch.empty((0,), dtype=torch.int64),
host_indices=torch.empty((0,), dtype=torch.int64),
)
)
class FakeEvictionStrategy:
def get_priority(self, node):
return 0
@@ -286,6 +299,41 @@ class TestHiRadixCacheCPBackup(CustomTestCase):
self.assertNotIn(1, root.children)
self.assertEqual(node.host_len, 16)
def test_write_backup_cp_success_returns_logical_length_for_zero_owned_rank(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = FakeZeroOwnedWriteController()
cache.ongoing_write_through = {}
cache.inc_node_lock_ref = lambda node: None
node = TreeNode()
node.id = 123
node.value = torch.arange(16, dtype=torch.int64)
backed_len = cache.write_backup(node)
self.assertEqual(backed_len, 16)
self.assertEqual(node.host_len, 16)
self.assertEqual(node.cp_hicache.host_indices.tolist(), [])
def test_attach_storage_backend_rejects_cp_hicache_without_controller_call(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.cache_controller = type(
"Controller",
(),
{
"attach_storage_backend": lambda *args, **kwargs: (_ for _ in ()).throw(
AssertionError("controller attach must not be called")
)
},
)()
ok, message = cache.attach_storage_backend("mooncake")
self.assertFalse(ok)
self.assertIn("CP shared KV", message)
self.assertIn("storage backend", message)
def test_evict_demotes_cp_backed_node_without_deleting_radix_child(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
@@ -515,6 +563,56 @@ class TestHiRadixCacheCPSplitEvict(CustomTestCase):
self.assertIsNotNone(parent.cp_hicache)
self.assertEqual(parent.cp_hicache.host_indices.tolist(), [80])
def test_synchronized_cp_host_eviction_removes_zero_owned_logical_leaf(self):
cache = HiRadixCache.__new__(HiRadixCache)
cache._uses_cp_hicache = True
cache.tp_world_size = 2
cache.tp_group = None
cache.root_node = TreeNode()
cache.root_node.key = RadixKey([])
cache.evictable_host_leaves = set()
cache.get_child_key_fn = lambda key: key.token_ids[0]
cache.eviction_strategy = type(
"Strategy", (), {"get_priority": lambda self, node: 0}
)()
cache._clear_pin = lambda node: None
cache._record_remove_event = lambda node: None
freed = []
cache.cache_controller = type(
"Controller",
(),
{
"evict_host": lambda self, indices: freed.append(indices.clone())
or len(indices)
},
)()
node = TreeNode()
node.parent = cache.root_node
node.key = RadixKey([1, 2, 3, 4])
node.value = None
node.host_len = 4
node.cp_hicache = CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.empty((0,), dtype=torch.int64),
host_indices=torch.empty((0,), dtype=torch.int64),
)
cache.root_node.children[1] = node
cache.evictable_host_leaves.add(node)
def mark_not_done(done_tensor, op=None, group=None):
done_tensor.fill_(0)
with patch("torch.distributed.all_reduce", side_effect=mark_not_done):
physical_freed = cache._evict_host_for_physical_slots(
0, synchronize_across_ranks=True
)
self.assertEqual(physical_freed, 0)
self.assertEqual(freed, [])
self.assertNotIn(1, cache.root_node.children)
self.assertEqual(node.host_len, 0)
self.assertIsNone(node.cp_hicache)
class TestHiRadixCacheCPLoadBack(CustomTestCase):
def test_cp_load_back_uses_host_len_not_host_value(self):