[UnifiedRadixCache] Fix HiCache load back start node (#25088)
This commit is contained in:
@@ -33,6 +33,7 @@ class _StubReq:
|
||||
self.prefix_indices = None
|
||||
self.last_node = None
|
||||
self.last_host_node = None
|
||||
self.best_match_node = None
|
||||
self.host_hit_length = None
|
||||
self.mamba_branching_seqlen = None
|
||||
self.cache_protected_len = None
|
||||
@@ -50,6 +51,7 @@ class TestZeroMatchResult(unittest.TestCase):
|
||||
self.assertEqual(int(zeroed.device_indices.numel()), 0)
|
||||
self.assertIs(zeroed.last_device_node, tree.root_node)
|
||||
self.assertIs(zeroed.last_host_node, tree.root_node)
|
||||
self.assertIs(zeroed.best_match_node, tree.root_node)
|
||||
self.assertEqual(zeroed.host_hit_length, 0)
|
||||
# dtype/device preserved (slice-not-allocate).
|
||||
self.assertEqual(zeroed.device_indices.dtype, match.device_indices.dtype)
|
||||
@@ -64,6 +66,7 @@ class TestZeroMatchResult(unittest.TestCase):
|
||||
device_indices=torch.empty((0,), dtype=torch.int64),
|
||||
last_device_node=None,
|
||||
last_host_node=None,
|
||||
best_match_node=None,
|
||||
host_hit_length=0,
|
||||
)
|
||||
self.assertIs(zero_match_result(_StubChunkCache(), original), original)
|
||||
|
||||
@@ -104,6 +104,7 @@ def test_preabort_detaches_session_and_preserves_slot():
|
||||
device_indices=torch.tensor([], dtype=torch.int64),
|
||||
last_device_node=None,
|
||||
last_host_node=None,
|
||||
best_match_node=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@@ -922,19 +922,19 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertEqual(node_count_before, 2)
|
||||
|
||||
tree._match_prefix_helper(RadixKey([1, 2]))
|
||||
value, best_node, best_value_len = tree._match_prefix_helper(
|
||||
value, best_match_node, best_value_len = tree._match_prefix_helper(
|
||||
RadixKey([1, 2, 3, 4])
|
||||
)
|
||||
self.assertEqual(best_value_len, 2)
|
||||
self.assertEqual(best_node.key.token_ids, [3, 4])
|
||||
self.assertEqual(best_match_node.key.token_ids, [3, 4])
|
||||
node_count_after_regular = count_nodes(tree.root_node)
|
||||
self.assertEqual(node_count_after_regular, node_count_before + 2)
|
||||
|
||||
value, best_node, best_value_len = tree._match_prefix_helper_readonly(
|
||||
value, best_match_node, best_value_len = tree._match_prefix_helper_readonly(
|
||||
RadixKey([1, 2, 3])
|
||||
)
|
||||
self.assertEqual(best_value_len, 1)
|
||||
self.assertEqual(best_node.key.token_ids, [1, 2])
|
||||
self.assertEqual(best_match_node.key.token_ids, [1, 2])
|
||||
node_count_after_readonly = count_nodes(tree.root_node)
|
||||
self.assertEqual(node_count_after_readonly, node_count_after_regular)
|
||||
|
||||
@@ -1821,6 +1821,7 @@ class UnifiedRadixCacheSuite:
|
||||
),
|
||||
last_device_node=leaf,
|
||||
last_host_node=leaf,
|
||||
best_match_node=leaf,
|
||||
host_hit_length=0,
|
||||
)
|
||||
result = swa_comp.finalize_match_result(
|
||||
@@ -1905,6 +1906,102 @@ class UnifiedRadixCacheSuite:
|
||||
n_swa,
|
||||
)
|
||||
|
||||
def _swa_anchor_chain_tokens(self, num_pages: int) -> list[int]:
|
||||
"""Reproduce the token sequence used by _build_chain_pages."""
|
||||
tokens: list[int] = []
|
||||
for i in range(num_pages):
|
||||
tokens.extend(self._make_seq(1000 * (i + 1), 1))
|
||||
return tokens
|
||||
|
||||
def _swa_anchor_setup(self):
|
||||
"""Chain layout (root-to-leaf):
|
||||
|
||||
... top padding ...
|
||||
chain[-(window_pages + 1)] = N : SWA tombstone
|
||||
chain[-(window_pages)] = Y : SWA host-only, FULL host-backed
|
||||
chain[-(window_pages - 1)..-2] : SWA device, FULL.host=None
|
||||
chain[-1] = X : SWA device, FULL.host=None
|
||||
|
||||
Anchor=X has window_pages of device pages below N, so the SWA
|
||||
load_back walker accumulates n_swa >= sliding_window_size and
|
||||
exits before reaching N. Anchor=Y has only 1 page and walks into
|
||||
N's tombstone.
|
||||
"""
|
||||
if not self.cfg.has_swa:
|
||||
self.skipTest("requires SWA")
|
||||
if self.cfg.has_mamba:
|
||||
self.skipTest("SWA-only path keeps the chain construction simple")
|
||||
ps = self.cfg.page_size
|
||||
sw = self.cfg.sliding_window_size
|
||||
if ps >= sw:
|
||||
self.skipTest("test scenario requires ps < sw")
|
||||
window_pages = (sw + ps - 1) // ps
|
||||
chain_pages = window_pages + 3
|
||||
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)
|
||||
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)
|
||||
|
||||
x = chain[-1]
|
||||
y = chain[-window_pages]
|
||||
n = chain[-(window_pages + 1)]
|
||||
|
||||
n.component_data[ComponentType.SWA].value = None
|
||||
n.component_data[ComponentType.SWA].host_value = None
|
||||
y.component_data[ComponentType.SWA].value = None
|
||||
# Strip FULL.host on X + intermediates so last_host_node walks past
|
||||
# them to Y. Y.FULL untouched preserves the leaf-up evict invariant.
|
||||
for node in chain[-(window_pages - 1) :]:
|
||||
node.component_data[ComponentType.FULL].host_value = None
|
||||
|
||||
tokens = self._swa_anchor_chain_tokens(len(chain))
|
||||
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()
|
||||
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_host_node, y)
|
||||
|
||||
def test_hicache_swa_load_back_anchored_on_best_match_node(self):
|
||||
tree, _, _, y, x, _ = self._swa_anchor_setup()
|
||||
ps = self.cfg.page_size
|
||||
swa_comp = tree.components[ComponentType.SWA]
|
||||
|
||||
transfers = swa_comp.build_hicache_transfers(x, CacheTransferPhase.LOAD_BACK)
|
||||
self.assertEqual(len(transfers), 1)
|
||||
xfer = transfers[0]
|
||||
self.assertEqual(xfer.name, PoolName.SWA)
|
||||
self.assertEqual(xfer.nodes_to_load, [y])
|
||||
self.assertEqual(int(xfer.host_indices.numel()), ps)
|
||||
|
||||
with self.assertRaises(AssertionError):
|
||||
swa_comp.build_hicache_transfers(y, CacheTransferPhase.LOAD_BACK)
|
||||
|
||||
def test_hicache_swa_finalize_anchored_on_best_match_node(self):
|
||||
tree, _, _, y, x, _ = self._swa_anchor_setup()
|
||||
swa_comp = tree.components[ComponentType.SWA]
|
||||
|
||||
base = MatchResult(
|
||||
device_indices=torch.empty((0,), dtype=torch.int64, device=tree.device),
|
||||
last_device_node=x,
|
||||
last_host_node=y,
|
||||
best_match_node=x,
|
||||
host_hit_length=0,
|
||||
)
|
||||
result = swa_comp.finalize_match_result(
|
||||
result=base,
|
||||
params=MatchPrefixParams(key=RadixKey(self._make_seq(1, 1))),
|
||||
value_chunks=[],
|
||||
best_value_len=0,
|
||||
)
|
||||
self.assertEqual(result.host_hit_length, 1)
|
||||
|
||||
def test_hicache_swa_temp_lock_does_not_release_restored_tombstone(self):
|
||||
"""A temporary scheduler lock that skipped a SWA tombstone must not
|
||||
release later load-back/request locks after the tombstone is restored.
|
||||
@@ -1950,6 +2047,61 @@ class UnifiedRadixCacheSuite:
|
||||
tree.dec_lock_ref(leaf, request_lock.to_dec_params())
|
||||
self.assertEqual(cd.lock_ref, 0)
|
||||
|
||||
def test_hicache_full_temp_lock_skips_evicted_anchor_and_mirrors_on_release(
|
||||
self,
|
||||
):
|
||||
"""Acquire records the evicted anchor in skip_lock_node_ids (phase 1)
|
||||
and locks device-on ancestors only (phase 2). After load_back
|
||||
restores the anchor, a second acquire covers it; releasing the
|
||||
first must mirror the skip so the anchor's lock_ref is not
|
||||
decremented twice.
|
||||
"""
|
||||
if self._skip_unsupported_hicache_test():
|
||||
return
|
||||
ps = self.cfg.page_size
|
||||
if 3 * ps > self.cfg.kv_size // 2:
|
||||
self.skipTest("kv_size too small")
|
||||
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")
|
||||
a, y, anchor = chain
|
||||
self._simulate_backup_tree(tree)
|
||||
|
||||
cd_anchor = anchor.component_data[ComponentType.FULL]
|
||||
cd_a = a.component_data[ComponentType.FULL]
|
||||
cd_y = y.component_data[ComponentType.FULL]
|
||||
anchor_value = cd_anchor.value
|
||||
cd_anchor.value = None
|
||||
|
||||
self.assertEqual(cd_anchor.lock_ref, 0)
|
||||
self.assertEqual(cd_y.lock_ref, 0)
|
||||
self.assertEqual(cd_a.lock_ref, 0)
|
||||
|
||||
temp_lock = tree.inc_lock_ref(anchor)
|
||||
self.assertEqual(cd_anchor.lock_ref, 0)
|
||||
self.assertEqual(cd_y.lock_ref, 1)
|
||||
self.assertEqual(cd_a.lock_ref, 1)
|
||||
self.assertIn(ComponentType.FULL, temp_lock.skip_lock_node_ids)
|
||||
self.assertIn(anchor.id, temp_lock.skip_lock_node_ids[ComponentType.FULL])
|
||||
|
||||
cd_anchor.value = anchor_value
|
||||
|
||||
second_lock = tree.inc_lock_ref(anchor)
|
||||
self.assertEqual(cd_anchor.lock_ref, 1)
|
||||
self.assertEqual(cd_y.lock_ref, 2)
|
||||
self.assertEqual(cd_a.lock_ref, 2)
|
||||
|
||||
tree.dec_lock_ref(anchor, temp_lock.to_dec_params())
|
||||
self.assertEqual(cd_anchor.lock_ref, 1)
|
||||
self.assertEqual(cd_y.lock_ref, 1)
|
||||
self.assertEqual(cd_a.lock_ref, 1)
|
||||
|
||||
tree.dec_lock_ref(anchor, second_lock.to_dec_params())
|
||||
self.assertEqual(cd_anchor.lock_ref, 0)
|
||||
self.assertEqual(cd_y.lock_ref, 0)
|
||||
self.assertEqual(cd_a.lock_ref, 0)
|
||||
|
||||
def test_hicache_mamba_temp_lock_does_not_release_restored_tombstone(self):
|
||||
"""A temporary scheduler lock that skipped a Mamba tombstone must not
|
||||
release later load-back/request locks after the tombstone is restored.
|
||||
|
||||
Reference in New Issue
Block a user