[UnifiedTree]: Fix UnifiedRadixCache device match semantics with HiCache (#25277)

This commit is contained in:
Zhangheng
2026-05-16 00:40:55 +08:00
committed by GitHub
parent 3f7e538b2f
commit 21b3ac52b4
7 changed files with 398 additions and 110 deletions
@@ -15,6 +15,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
EvictResult,
InitLoadBackParams,
InsertParams,
MatchPrefixParams,
MatchResult,
@@ -908,6 +909,8 @@ class UnifiedRadixCacheSuite:
"""Verify readonly match does not modify tree structure (no split)."""
if self.cfg.page_size > 1 or self.cfg.has_mamba or self.cfg.has_swa:
self.skipTest("Full-only page_size=1 only")
if not hasattr(UnifiedRadixCache, "_match_prefix_helper_readonly"):
self.skipTest("_match_prefix_helper_readonly is not available")
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
self._insert(tree, allocator, req_to_token_pool, [1, 2, 3, 4, 5])
@@ -922,19 +925,27 @@ class UnifiedRadixCacheSuite:
self.assertEqual(node_count_before, 2)
tree._match_prefix_helper(RadixKey([1, 2]))
value, best_match_node, best_value_len = tree._match_prefix_helper(
RadixKey([1, 2, 3, 4])
)
(
value,
best_match_node,
best_match_device_node,
best_value_len,
) = tree._match_prefix_helper(RadixKey([1, 2, 3, 4]))
self.assertEqual(best_value_len, 2)
self.assertEqual(best_match_node.key.token_ids, [3, 4])
self.assertIs(best_match_device_node, best_match_node)
node_count_after_regular = count_nodes(tree.root_node)
self.assertEqual(node_count_after_regular, node_count_before + 2)
value, best_match_node, best_value_len = tree._match_prefix_helper_readonly(
RadixKey([1, 2, 3])
)
(
value,
best_match_node,
best_match_device_node,
best_value_len,
) = tree._match_prefix_helper_readonly(RadixKey([1, 2, 3]))
self.assertEqual(best_value_len, 1)
self.assertEqual(best_match_node.key.token_ids, [1, 2])
self.assertIs(best_match_device_node, best_match_node)
node_count_after_readonly = count_nodes(tree.root_node)
self.assertEqual(node_count_after_readonly, node_count_after_regular)
@@ -1258,8 +1269,8 @@ class UnifiedRadixCacheSuite:
# ================================================================
def _skip_unsupported_hicache_test(self):
if self.cfg.has_swa:
self.skipTest("HiCache tests do not run on SWA stacks")
if self.cfg.has_swa and self.cfg.has_mamba:
self.skipTest("HiCache unit fixture does not support SWA + Mamba stacks")
return False
def _simulate_backup(self, tree, node):
@@ -1343,14 +1354,14 @@ class UnifiedRadixCacheSuite:
self._backup_node(tree, node)
def _load_back_node(self, tree, node):
device_indices = tree.load_back(node)
self.assertIsNotNone(device_indices)
loaded = tree.load_back(node)
self.assertTrue(loaded)
producer_id = tree.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(tree.cache_controller.ack_load_queue):
finish_event.synchronize()
tree.loading_check()
return device_indices
return node.component_data[ComponentType.FULL].value
def _get_full_kv_pool(self, allocator):
kv_pool = allocator.get_kvcache()
@@ -1474,7 +1485,9 @@ class UnifiedRadixCacheSuite:
def test_hicache_partial_match_splits_evicted_backed_up_node(self):
"""Partial matches on host-only nodes must keep the host prefix usable."""
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
if self._skip_unsupported_hicache_test():
return
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
ps = self.cfg.page_size
seq = self._make_seq(1, 4)
expected_prefix = seq[: 2 * ps]
@@ -1484,7 +1497,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
node = m.last_device_node
self._simulate_backup(tree, node)
self._backup_node(tree, node)
tree.evict(EvictParams(num_tokens=len(seq)))
self.assertTrue(node.evicted)
@@ -1684,6 +1697,244 @@ class UnifiedRadixCacheSuite:
chain.reverse()
return chain
def _release_ongoing_load_back_locks(self, tree):
for node, lock_params in list(tree.ongoing_load_back.values()):
tree.dec_lock_ref(node, lock_params)
tree.ongoing_load_back.clear()
def _finish_pending_loads(self, tree):
producer_id = tree.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(tree.cache_controller.ack_load_queue):
finish_event.synchronize()
tree.loading_check()
def _match_tokens_for_chain(self, chain):
tokens: list[int] = []
for node in chain:
tokens.extend(node.key.token_ids)
return tokens
def _set_aux_host_tombstone(self, tree, node, component_type):
cd = node.component_data[component_type]
self.assertIsNotNone(cd.value)
if cd.host_value is None:
cd.host_value = cd.value.clone()
old_value = cd.value
cd.value = None
if component_type in tree.lru_lists and tree.lru_lists[component_type].in_list(
node
):
tree.lru_lists[component_type].remove_node(node)
tree.host_lru_lists[component_type].insert_mru(node)
tree.component_evictable_size_[component_type] -= len(old_value)
def test_match_prefix_best_and_device_node_without_hicache(self):
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
ps = self.cfg.page_size
min_tokens = 2 * ps
if self.cfg.has_swa:
min_tokens = max(min_tokens, self.cfg.sliding_window_size + ps)
seq = self._make_seq(1, (min_tokens + ps - 1) // ps)
self._insert(tree, allocator, req_to_token_pool, seq)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
self.assertEqual(len(result.device_indices), len(seq))
self.assertIs(result.best_match_node, result.last_device_node)
self.assertIs(result.last_host_node, result.last_device_node)
self.assertEqual(result.host_hit_length, 0)
def test_hicache_mamba_host_best_match_keeps_device_anchor(self):
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3)
if len(chain) < 3:
self.skipTest("chain too short")
leaf = chain[-1]
parent = chain[-2]
tokens = self._match_tokens_for_chain(chain)
self._backup_node(tree, leaf)
tree.evict(EvictParams(num_tokens=len(leaf.key)))
self.assertTrue(leaf.evicted)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertIs(result.best_match_node, leaf)
self.assertIs(result.last_device_node, parent)
self.assertEqual(len(result.device_indices), len(tokens) - len(leaf.key))
self.assertEqual(result.host_hit_length, len(leaf.key))
def test_hicache_swa_host_best_match_keeps_device_anchor(self):
if not self.cfg.has_swa or self.cfg.has_mamba or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+SWA")
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3)
if len(chain) < 3:
self.skipTest("chain too short")
leaf = chain[-1]
parent = chain[-2]
tokens = self._match_tokens_for_chain(chain)
self._backup_node(tree, leaf)
tree.evict(EvictParams(num_tokens=len(leaf.key)))
self.assertTrue(leaf.evicted)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertIs(result.best_match_node, leaf)
self.assertIs(result.last_device_node, parent)
self.assertEqual(len(result.device_indices), len(tokens) - len(leaf.key))
self.assertEqual(result.host_hit_length, 1)
def test_mamba_branching_seqlen_disabled_under_hicache(self):
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
chunk_size = get_global_server_args().mamba_cache_chunk_size
tokens = self._make_seq(1, chunk_size + 1)
self._insert(tree, allocator, req_to_token_pool, tokens)
leaf = tree.match_prefix(
MatchPrefixParams(key=RadixKey(tokens))
).last_device_node
mamba_cd = leaf.component_data[ComponentType.MAMBA]
mamba_cd.value = None
no_hicache = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertIs(no_hicache.best_match_node, tree.root_node)
self.assertIs(no_hicache.last_device_node, tree.root_node)
self.assertEqual(no_hicache.mamba_branching_seqlen, chunk_size)
tree_h, allocator_h, req_to_token_pool_h = self._build_hicache_fixture()
self._insert(tree_h, allocator_h, req_to_token_pool_h, tokens)
leaf_h = tree_h.match_prefix(
MatchPrefixParams(key=RadixKey(tokens))
).last_device_node
self._backup_node(tree_h, leaf_h)
tree_h.evict(EvictParams(num_tokens=len(tokens)))
with_hicache = tree_h.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertIs(with_hicache.best_match_node, leaf_h)
self.assertIs(with_hicache.last_device_node, tree_h.root_node)
self.assertIsNone(with_hicache.mamba_branching_seqlen)
def test_scheduler_hicache_full_mamba_init_load_back_appends_new_indices(self):
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3)
if len(chain) < 3:
self.skipTest("chain too short")
leaf = chain[-1]
tokens = self._match_tokens_for_chain(chain)
self._backup_node(tree, leaf)
tree.evict(EvictParams(num_tokens=len(leaf.key)))
self.assertTrue(leaf.evicted)
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
req.host_hit_length = match.host_hit_length
new_indices, new_node = tree.init_load_back(
InitLoadBackParams(
best_match_node=req.best_match_node,
host_hit_length=req.host_hit_length,
req=req,
)
)
self.assertIs(new_node, leaf)
self.assertEqual(len(torch.cat([req.prefix_indices, new_indices])), len(tokens))
self.assertIsNotNone(leaf.component_data[ComponentType.MAMBA].value)
self._finish_pending_loads(tree)
self._release_ongoing_load_back_locks(tree)
def test_scheduler_hicache_aux_only_load_back_appends_full_device_indices(self):
if self.cfg.page_size != 1:
self.skipTest("page_size=1 keeps the expected suffix precise")
aux = None
if self.cfg.has_swa and not self.cfg.has_mamba:
aux = ComponentType.SWA
elif self.cfg.has_mamba and not self.cfg.has_swa:
aux = ComponentType.MAMBA
if aux is None:
self.skipTest("requires exactly one aux component")
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3)
if len(chain) < 3:
self.skipTest("chain too short")
leaf = chain[-1]
tokens = self._match_tokens_for_chain(chain)
leaf_full = leaf.component_data[ComponentType.FULL].value.clone()
self._backup_node(tree, leaf)
self._set_aux_host_tombstone(tree, leaf, aux)
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
req.host_hit_length = match.host_hit_length
new_indices, new_node = tree.init_load_back(
InitLoadBackParams(
best_match_node=req.best_match_node,
host_hit_length=req.host_hit_length,
req=req,
)
)
self.assertIs(new_node, leaf)
self.assertEqual(new_indices.tolist(), leaf_full.tolist())
self.assertEqual(len(torch.cat([req.prefix_indices, new_indices])), len(tokens))
self.assertEqual(
leaf.component_data[ComponentType.FULL].value.tolist(),
leaf_full.tolist(),
)
self.assertIsNotNone(leaf.component_data[aux].value)
self._finish_pending_loads(tree)
self._release_ongoing_load_back_locks(tree)
def test_scheduler_hicache_load_back_fallback_keeps_old_anchor(self):
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3)
if len(chain) < 3:
self.skipTest("chain too short")
leaf = chain[-1]
tokens = self._match_tokens_for_chain(chain)
self._backup_node(tree, leaf)
tree.evict(EvictParams(num_tokens=len(leaf.key)))
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
req.host_hit_length = match.host_hit_length
new_indices, new_node = tree.init_load_back(
InitLoadBackParams(
best_match_node=req.best_match_node,
host_hit_length=req.host_hit_length,
req=req,
mem_quota=-1_000_000,
)
)
self.assertEqual(len(new_indices), 0)
self.assertIs(new_node, match.last_device_node)
self.assertIsNone(leaf.component_data[ComponentType.FULL].value)
self.assertIsNone(leaf.component_data[ComponentType.MAMBA].value)
def test_hicache_swa_load_back_min_suffix(self):
"""LOAD_BACK collects only the suffix nodes needed to cover sliding_window_size."""
if not self.cfg.has_swa:
@@ -1940,11 +2191,11 @@ class UnifiedRadixCacheSuite:
if chain_pages * ps > self.cfg.kv_size // 2:
self.skipTest("kv_size too small for the desired chain")
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
tree, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(tree, allocator, req_to_token_pool, chain_pages)
if len(chain) < chain_pages:
self.skipTest("chain too short")
self._simulate_backup_tree(tree)
self._backup_tree(tree)
x = chain[-1]
y = chain[-window_pages]
@@ -1962,10 +2213,10 @@ class UnifiedRadixCacheSuite:
return tree, chain, n, y, x, tokens
def test_hicache_swa_match_prefix_picks_best_match_node_above_last_host(self):
tree, _, _, y, x, tokens = self._swa_anchor_setup()
tree, _, n, y, x, tokens = self._swa_anchor_setup()
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertIs(result.best_match_node, x)
self.assertIs(result.last_device_node, x)
self.assertIs(result.last_device_node, n.parent)
self.assertIs(result.last_host_node, y)
def test_hicache_swa_load_back_anchored_on_best_match_node(self):