[Unified Tree] Port SWA Branching-Point Caching to the Rust TreeCore (#37584)

This commit is contained in:
Shuwen Wang
2026-09-12 18:37:57 +08:00
committed by GitHub
parent 0b415fa573
commit bd45cd50ca
15 changed files with 762 additions and 109 deletions
@@ -101,6 +101,7 @@ def _pump_insert(core: RustUnifiedTreeCore, params: InsertParams) -> InsertResul
prefix_len=step.result.prefix_len,
last_device_node=step.result.last_device_node,
mamba_exist=step.result.mamba_exist,
swa_branch_inserted=step.result.swa_branch_inserted,
cache_actions=actions,
)
@@ -2288,5 +2289,29 @@ def test_stale_inspection_handles_raise_key_error_or_report_absence():
assert not core.is_node_in_host_lru(stale_root, ComponentType.SWA)
# ---- SWA branching-point caching ----
def _swa_hicache_core(window: int = 8) -> RustUnifiedTreeCore:
core = _swa_tree_core(window=window)
core.set_hicache_enabled()
core.has_swa_host_pool = True
return core
def test_insert_reports_whether_it_reached_the_swa_branch_boundary():
for branching_seqlen, expected in [(2, True), (3, False), (None, False)]:
core = _swa_hicache_core()
result = _pump_insert(
core,
InsertParams(
key=_key([1, 2]),
value=torch.tensor([10, 11], dtype=torch.int64),
swa_branching_seqlen=branching_seqlen,
),
)
assert result.swa_branch_inserted is expected, branching_seqlen
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -5628,6 +5628,44 @@ class UnifiedRadixCacheSuite:
cache.tree_core.has_swa_host_pool, swa._swa_kv_pool_host is not None
)
def test_swa_backup_collector_is_shared_by_both_call_sites(self):
"""needs_incremental_backup and the BACKUP_HOST transfer read one
collector: cache mode walks the window past a host-backed target to
its device-only ancestor, buffer mode stages the target alone."""
if not self.cfg.has_swa or self.cfg.has_mamba:
self.skipTest("requires SWA-only")
if self.cfg.sliding_window_size <= self.cfg.page_size:
self.skipTest("the window must reach past the leaf's own page")
if _selected_tree_core_test_backend() == "rust":
# needs_incremental_backup is a component method on Python nodes;
# the Rust core pins the same contract in its own unit suite.
self.skipTest("component-level check is Python-core only")
cache, allocator, req_to_token_pool = self._build_hicache_fixture()
chain = self._build_chain_pages(cache, allocator, req_to_token_pool, 2)
if len(chain) < 2:
self.skipTest("chain too short")
parent, leaf = chain[-2], chain[-1]
swa = cache.components[ComponentType.SWA]
leaf_node = cache.tree_core.node_by_id(leaf)
cache.tree_core.set_component_host_value_raw(
leaf,
ComponentType.SWA,
_device_value(cache, leaf, ComponentType.SWA).clone(),
)
self.assertTrue(swa.needs_incremental_backup(leaf_node))
xfer = cache.tree_core.build_hicache_transfers(
ComponentType.SWA, leaf, CacheTransferPhase.BACKUP_HOST
)[0]
self.assertEqual(xfer.nodes_to_load, [parent])
cache.tree_core.set_host_memory_buffer_only()
self.assertTrue(swa.needs_incremental_backup(leaf_node))
xfer = cache.tree_core.build_hicache_transfers(
ComponentType.SWA, leaf, CacheTransferPhase.BACKUP_HOST
)[0]
self.assertEqual(xfer.nodes_to_load, [leaf])
def test_zero_match_result_carries_node_id_handles(self):
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
ps = self.cfg.page_size
@@ -5872,14 +5910,7 @@ class UnifiedRadixCacheSuite:
self.assertEqual(result.host_hit_length, 0)
self.assertEqual(result.swa_host_hit_length, _node_key_length(cache, leaf))
def _skip_swa_branching_on_rust(self) -> None:
# TODO(alphabetc1): drop this gate once #37584 ports SWA branching-point
# caching to the Rust tree core.
if _selected_tree_core_test_backend() == "rust":
self.skipTest("SWA branching-point caching is Python-core only")
def test_swa_branching_seqlen_uses_device_full_hit(self):
self._skip_swa_branching_on_rust()
if (
not self.cfg.has_swa
or self.cfg.has_mamba
@@ -5895,20 +5926,13 @@ class UnifiedRadixCacheSuite:
self._insert(cache, allocator, req_to_token_pool, prefix)
self._insert(cache, allocator, req_to_token_pool, tokens)
leaf = cache.resolve_node_handle(
cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
leaf = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
evicted = cache.tree_core.evict_component(
leaf, ComponentType.SWA, EvictLayer.DEVICE
)
device_frees = defaultdict(list)
cache.tree_core._evict_component_and_detach_lru(
leaf,
cache.components[ComponentType.SWA],
device_frees=device_frees,
host_frees=defaultdict(list),
target=EvictLayer.DEVICE,
)
cache._drain_device_frees(device_frees)
cache._free_values(evicted.device_frees, evicted.host_frees)
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
@@ -5931,7 +5955,6 @@ class UnifiedRadixCacheSuite:
self.assertIsNone(rematch.swa_branching_seqlen)
def test_swa_branching_seqlen_uses_host_full_hit(self):
self._skip_swa_branching_on_rust()
if (
not self.cfg.has_swa
or self.cfg.has_mamba
@@ -5946,24 +5969,21 @@ class UnifiedRadixCacheSuite:
self._insert(cache, allocator, req_to_token_pool, prefix)
self._insert(cache, allocator, req_to_token_pool, tokens)
leaf = cache.resolve_node_handle(
cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
)
parent = leaf.parent
self._backup_node(cache, leaf.id)
lock_result = cache.inc_lock_ref(parent.id)
leaf = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
leaf_len = _node_key_length(cache, leaf)
parent = _node_parent(cache, leaf)
self._backup_node(cache, leaf)
lock_result = cache.inc_lock_ref(parent)
try:
cache.evict(EvictParams(num_tokens=len(leaf.key)))
cache.evict(EvictParams(num_tokens=leaf_len))
finally:
cache.dec_lock_ref(parent.id, lock_result.to_dec_params())
device_frees = defaultdict(list)
host_frees = defaultdict(list)
cache.components[ComponentType.SWA].evict_component(
leaf, device_frees, host_frees, target=EvictLayer.HOST
cache.dec_lock_ref(parent, lock_result.to_dec_params())
evicted = cache.tree_core.evict_component(
leaf, ComponentType.SWA, EvictLayer.HOST
)
cache._free_values(device_frees, host_frees)
cache._free_values(evicted.device_frees, evicted.host_frees)
full_host_pool = cache.cache_controller.mem_pool_host
swa_host_pool = cache.components[ComponentType.SWA]._swa_kv_pool_host
full_available_before = full_host_pool.available_size()
@@ -5985,7 +6005,7 @@ class UnifiedRadixCacheSuite:
self.assertEqual(full_host_pool.available_size(), full_available_before)
self.assertEqual(
swa_host_pool.available_size(),
swa_available_before - len(leaf.key),
swa_available_before - leaf_len,
)
rematch = cache.match_prefix(
@@ -6799,7 +6819,6 @@ class UnifiedRadixCacheSuite:
self.assertEqual(comp_xfers[ComponentType.SWA][0].nodes_to_load, [a, b])
def test_hicache_swa_backup_window_stops_at_pending_ancestor(self):
self._skip_swa_branching_on_rust()
if (
not self.cfg.has_swa
or self.cfg.has_mamba
@@ -6816,24 +6835,30 @@ class UnifiedRadixCacheSuite:
c = chain[-1]
c_swa = _device_value(cache, c, ComponentType.SWA).clone()
# First transfer: publish Full for C only, leaving SWA dirty while the
# write-through ack is still pending.
# First transfer: publish Full for C only. C's SWA is a device tombstone
# (decode-evicted, never backed up), so the write-through ack stays
# pending on a node the SWA backup window has nothing to send for.
cache.tree_core.set_component_device_value_raw(c, ComponentType.SWA, None)
if cache.tree_core.is_node_in_device_lru(c, ComponentType.SWA):
cache.tree_core.remove_node_from_device_lru(c, ComponentType.SWA)
cache.tree_core.set_component_evictable_size(
ComponentType.SWA,
cache.tree_core.component_evictable_size(ComponentType.SWA) - len(c_swa),
)
self.assertGreater(
cache._execute_and_commit_kv_backup(BackupKV(node_ids=[c])),
0,
)
self.assertEqual(
cache.tree_core.node_by_id(c).write_through_pending_id,
c,
)
self.assertEqual(cache.tree_core.get_write_through_pending_id(c), c)
self.assertIsNotNone(_host_value(cache, c, ComponentType.FULL))
self.assertIsNone(_host_value(cache, c, ComponentType.SWA))
# Simulate SWA being reconstructed on device before the first ack. The
# next incremental SWA backup must treat C as the boundary and back up
# only the newly inserted descendant.
cache.tree_core.set_component_device_value_raw(c, ComponentType.SWA, c_swa)
# SWA is reconstructed on device before the first ack, the way a
# load-back commit stores it: under the pending segment lock the value
# counts as protected until the ack releases it. The next incremental
# SWA backup must treat C as the boundary and back up only the newly
# inserted descendant.
cache.tree_core.set_component_device_value(c, ComponentType.SWA, c_swa)
tokens = self._match_tokens_for_chain(cache, chain)
next_tokens = tokens + self._make_seq(9000, 1)
cache.write_through_threshold = 1
@@ -6843,20 +6868,15 @@ class UnifiedRadixCacheSuite:
d = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", next_tokens)))
).last_device_node
self.assertEqual(
cache.tree_core.node_by_id(c).write_through_pending_id,
c,
)
self.assertEqual(
cache.tree_core.node_by_id(d).write_through_pending_id,
d,
)
self.assertEqual(cache.tree_core.get_write_through_pending_id(c), c)
self.assertEqual(cache.tree_core.get_write_through_pending_id(d), d)
self.assertIsNone(_host_value(cache, c, ComponentType.SWA))
self.assertIsNotNone(_host_value(cache, d, ComponentType.SWA))
cache.writing_check(write_back=True)
self.assertIsNone(cache.tree_core.node_by_id(c).write_through_pending_id)
self.assertIsNone(cache.tree_core.node_by_id(d).write_through_pending_id)
self.assertIsNone(cache.tree_core.get_write_through_pending_id(c))
self.assertIsNone(cache.tree_core.get_write_through_pending_id(d))
cache.sanity_check()
def _swa_finalize_setup(self):
"""Build a SWA chain long enough to fill at least the window