diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index db80063997a3..2251efb4d24f 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -691,6 +691,11 @@ class Envs: # - Source builds with a missing or unusable Rust toolchain. # This also applies when Rust is explicitly selected. SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("rust") + # Python TreeCore only: keep a persistent lazy-deletion heap over the Full + # component's evictable leaves instead of rebuilding it on every eviction + # call. False re-keys every leaf at each eviction (legacy O(#leaves) cost, + # identical eviction order) through the same code path. + SGLANG_UNIFIED_RADIX_LAZY_EVICTION_HEAP = EnvBool(True) # Once decode passes the sliding window, drop the SWA part of the prefill's # tree lock; its SWA KV becomes evictable rather than freed. SGLANG_OPT_RELEASE_PREFILL_SWA = EnvBoolWithAlias( diff --git a/python/sglang/srt/mem_cache/unified_cache/components/full.py b/python/sglang/srt/mem_cache/unified_cache/components/full.py index f9dc55c82fa2..c2ae4376b5ab 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/full.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/full.py @@ -1,6 +1,5 @@ from __future__ import annotations -import heapq from typing import TYPE_CHECKING, Callable, Optional, Sequence import torch @@ -67,6 +66,7 @@ def _dec_session_coverage(self, session_id: str, leaf: UnifiedTreeNode) -> None: cd = node.component_data[self.component_type] assert cd.session_ref > 0 cd.session_ref -= 1 + self.tree_core._touch_full_eviction_key(node) node = node.parent def _advance_session_coverage( @@ -83,6 +83,7 @@ def _advance_session_coverage( and node is not self.tree_core.root_node ): node.component_data[self.component_type].session_ref += 1 + self.tree_core._touch_full_eviction_key(node) node = node.parent def _recede_session_coverage( @@ -101,6 +102,7 @@ def _recede_session_coverage( cd = node.component_data[self.component_type] assert cd.session_ref > 0 cd.session_ref -= 1 + self.tree_core._touch_full_eviction_key(node) node = node.parent def create_match_validator( @@ -200,14 +202,9 @@ def _session_ref_eviction_strategy(self, node: UnifiedTreeNode): return ref > 0, ref, self.tree_core.eviction_strategy.get_priority(node) def _evict_device_start(self, request_cnt: int) -> None: - self._ensure_eviction_strategy() self._evict_device_request_cnt = request_cnt self._evict_device_last_node = None - self._evict_device_heap = [ - (self.session_ref_eviction_strategy(n), n) - for n in self.tree_core.evictable_device_leaves - ] - heapq.heapify(self._evict_device_heap) + self.tree_core.full_device_heap.begin_walk() def _evict_device_next_node( self, @@ -216,28 +213,23 @@ def _evict_device_next_node( host_frees: dict[ComponentType, list[torch.Tensor]], ) -> Optional[NodeId]: ct = self.component_type + heap = self.tree_core.full_device_heap lv = self._evict_device_last_node - if ( - lv is not None - and lv.parent is not None - and lv.parent in self.tree_core.evictable_device_leaves - ): - heapq.heappush( - self._evict_device_heap, - (self.session_ref_eviction_strategy(lv.parent), lv.parent), - ) + if lv is not None and lv.parent is not None: + # The evicted leaf's parent may have become a device leaf: admit + # it to the current walk, keyed now (the legacy explicit push). + heap.promote(lv.parent) self._evict_device_last_node = None - while tracker[ct] < self._evict_device_request_cnt and self._evict_device_heap: - _, x = heapq.heappop(self._evict_device_heap) - if x not in self.tree_core.evictable_device_leaves: - continue - self._evict_device_last_node = x - return x.id + if tracker[ct] < self._evict_device_request_cnt: + x = heap.pop_next() + if x is not None: + self._evict_device_last_node = x + return x.id return None def _evict_device_end(self) -> None: - self._evict_device_heap = [] self._evict_device_last_node = None + self.tree_core.full_device_heap.end_walk() def drive_host_eviction( self, @@ -247,26 +239,19 @@ def drive_host_eviction( host_frees: dict[ComponentType, list[torch.Tensor]], ) -> None: """Evict host leaves to free KV host pool space.""" - self._ensure_eviction_strategy() - heap = [ - (self.session_ref_eviction_strategy(n), n) - for n in self.tree_core.evictable_host_leaves - ] - heapq.heapify(heap) ct = self.component_type - while tracker[ct] < num_tokens and heap: - _, x = heapq.heappop(heap) - if x not in self.tree_core.evictable_host_leaves: - continue - self.tree_core._evict_host_leaf(x, tracker, device_frees, host_frees) - if ( - x.parent is not None - and x.parent in self.tree_core.evictable_host_leaves - ): - heapq.heappush( - heap, - (self.session_ref_eviction_strategy(x.parent), x.parent), - ) + heap = self.tree_core.full_host_heap + heap.begin_walk() + try: + while tracker[ct] < num_tokens: + x = heap.pop_next() + if x is None: + break + self.tree_core._evict_host_leaf(x, tracker, device_frees, host_frees) + if x.parent is not None: + heap.promote(x.parent) + finally: + heap.end_walk() def acquire_component_lock( self, @@ -301,6 +286,9 @@ def acquire_component_lock( # Lock the device-on segment up to root delta = 0 + tree_core = self.tree_core + evictable_device_leaves = tree_core.evictable_device_leaves + device_live = tree_core._full_device_live while cur is not root: cd = cur.component_data[ct] assert cd.value is not None, ( @@ -308,11 +296,13 @@ def acquire_component_lock( ) if cd.lock_ref == 0: key_len = len(cd.value) - self.tree_core.component_evictable_size_[ct] -= key_len - self.tree_core.component_protected_size_[ct] += key_len + tree_core.component_evictable_size_[ct] -= key_len + tree_core.component_protected_size_[ct] += key_len delta += key_len cd.lock_ref += 1 - self.tree_core.evictable_device_leaves.discard(cur) + evictable_device_leaves.discard(cur) + if cur in device_live: + tree_core.full_device_heap.forget(cur) cur = cur.parent result.delta = delta return result diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba.py index 14e466917f3d..3af549e2670b 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba.py @@ -241,10 +241,12 @@ def commit_insert_component_data( node.id, self.component_type, params.mamba_value ) node.last_access_time = get_and_increase_time_counter() + self.tree_core._touch_full_eviction_key(node) self._emit_excess_path_states_eviction(node, cache_actions) return self.tree_core.lru_lists[self.component_type].reset_node_mru(node) node.last_access_time = get_and_increase_time_counter() + self.tree_core._touch_full_eviction_key(node) result.mamba_exist = True def _emit_excess_path_states_eviction( diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py index 93d640ed101f..5effd533c290 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py @@ -16,6 +16,7 @@ from __future__ import annotations +import heapq import logging import sys from array import array @@ -389,6 +390,189 @@ def get_lru_no_host_lock(self): # WALK (one node per step) -> COMMIT (leaf + commit hooks) -> TAIL (refresh + backup). +class _LazyLeafHeap: + """Persistent min-heap over one evictable-leaf set with lazy invalidation. + + Replaces rebuilding ``[(key(n), n) for n in leaves]`` + ``heapify`` on every + eviction call (O(#leaves)) with a heap that survives across calls: + + * ``_live`` maps every member of ``members`` to the key it was last pushed + with; a heap entry ``(key, node)`` is *valid* iff ``_live[node] == key``, + everything else is stale and skipped on pop (lazy deletion). + * ``refresh``/``touch`` re-push a node when its key input changed, + ``forget`` drops it when it leaves the set, ``promote`` is the walk-time + parent push. Stale entries are bounded by compaction (``_compact``). + * A walk (``begin_walk``/``pop_next``/``end_walk``) sees exactly what the + rebuilt heap used to see: the keys frozen at ``begin_walk``, nodes that + enter the set during the walk stay invisible (parked in ``_pending``) + unless explicitly ``promote``d, and a yielded node loses its live entry + until ``end_walk`` re-keys it if it is still a member. Eviction order is + therefore identical to the per-call rebuild for every eviction strategy. + + Entries keep the legacy ``(key, node)`` shape so ``UnifiedTreeNode.__lt__`` + remains the tie-break. The key function is resolved lazily on first use. + """ + + __slots__ = ( + "_members", + "_key_provider", + "_key_fn", + "_heap", + "_live", + "_pending", + "_yielded", + "rebuild_each_walk", + ) + + def __init__( + self, + members: set, + key_provider: Callable[[], Callable[[Any], Any]], + rebuild_each_walk: bool = False, + ) -> None: + self._members = members + self._key_provider = key_provider + self._key_fn: Optional[Callable[[Any], Any]] = None + self._heap: list = [] + self._live: dict = {} + # ``None`` outside a walk; lists while a walk is in progress. + self._pending: Optional[list] = None + self._yielded: Optional[list] = None + # Kill switch: re-key every member at ``begin_walk`` (legacy cost, + # identical order) through the same code path. + self.rebuild_each_walk = rebuild_each_walk + + def _key(self, node): + key_fn = self._key_fn + if key_fn is None: + key_fn = self._key_fn = self._key_provider() + return key_fn(node) + + @property + def walking(self) -> bool: + return self._pending is not None + + def __len__(self) -> int: + return len(self._live) + + def refresh(self, node) -> None: + """Membership hook: (re-)key ``node`` if it is a member of the set.""" + if node not in self._members: + return + if self._pending is not None: + self._pending.append(node) + return + self._upsert(node) + + def touch(self, node) -> None: + """Key-mutation hook: one dict lookup for non-members.""" + if node in self._live: + self.refresh(node) + + def promote(self, node) -> None: + """Walk-time explicit push (the evicted leaf's parent), visible now.""" + if node in self._members: + self._upsert(node) + + def forget(self, node) -> None: + """Membership-removal hook.""" + if self._live.pop(node, None) is not None: + self._maybe_compact() + + def _upsert(self, node) -> None: + key = self._key(node) + live = self._live + if live.get(node) == key: + return + live[node] = key + heapq.heappush(self._heap, (key, node)) + self._maybe_compact() + + def begin_walk(self) -> None: + assert self._pending is None, "eviction walk already in progress" + if self.rebuild_each_walk: + for node in self._members: + self._upsert(node) + self._pending = [] + self._yielded = [] + + def pop_next(self): + """Next valid victim in key order, or ``None`` when exhausted.""" + heap = self._heap + live = self._live + members = self._members + while heap: + key, node = heapq.heappop(heap) + if live.get(node) != key: + continue # superseded, forgotten or already yielded + del live[node] + if node not in members: + continue # defensive: should have been forgotten + self._yielded.append(node) + return node + return None + + def end_walk(self) -> None: + pending, yielded = self._pending, self._yielded + self._pending = None + self._yielded = None + # Declined victims regain a live entry; entrants born during the walk + # and members touched mid-walk get their fresh key. + for node in yielded: + self.refresh(node) + for node in pending: + self.refresh(node) + # A yielded victim that the caller destroyed drops its live entry + # without passing a bound check; a walk leaves one stale entry per evicted leaf. + self._maybe_compact() + + def _compact_threshold(self) -> int: + return 2 * len(self._live) + 64 + + def _maybe_compact(self) -> None: + if len(self._heap) > self._compact_threshold(): + self._compact() + + def _compact(self) -> None: + # Filter ``_live`` in place: the tree core keeps direct references to + # it for its inlined membership checks. + live = self._live + members = self._members + for node in [n for n in live if n not in members]: + del live[node] + self._heap = [(k, n) for n, k in live.items()] + heapq.heapify(self._heap) + + def check_invariants(self, report: Callable[[str], Any], name: str) -> None: + """Append invariant violations via ``report``; only meaningful + outside a walk (a walk is allowed to have yielded members).""" + if self._pending is not None: + return + members = self._members + live = self._live + extra = [n for n in live if n not in members] + missing = [n for n in members if n not in live] + if extra: + report(f"[{name}] live but not member: {[n.id for n in extra[:5]]}") + if missing: + report(f"[{name}] member without live entry: {[n.id for n in missing[:5]]}") + stale_key = [n for n in members if n in live and live[n] != self._key(n)] + if stale_key: + report(f"[{name}] live key out of date: {[n.id for n in stale_key[:5]]}") + heap_keys: dict = {} + for key, node in self._heap: + heap_keys.setdefault(node, []).append(key) + no_entry = [n for n, k in live.items() if k not in heap_keys.get(n, ())] + if no_entry: + report( + f"[{name}] live entry missing from heap: {[n.id for n in no_entry[:5]]}" + ) + if len(self._heap) > self._compact_threshold(): + report( + f"[{name}] heap not compacted: {len(self._heap)} entries for {len(live)} live" + ) + + class _InsertPhase(Enum): WALK = auto() COMMIT = auto() @@ -472,6 +656,26 @@ def _session_lru_predicate(self, ct: ComponentType): return None return lambda node: node.component_data[ct].session_ref > 0 + def _full_eviction_key_fn(self) -> Callable[[UnifiedTreeNode], Any]: + """Eviction key of the Full component (session-ref tuple when session + radix cache is on, else the strategy priority). Bound lazily because + the component binds its strategy only after the tree attaches.""" + full = self.components_by_type[BASE_COMPONENT_TYPE] + ensure = getattr(full, "_ensure_eviction_strategy", None) + if ensure is not None: + ensure() + key_fn = getattr(full, "session_ref_eviction_strategy", None) + return key_fn if key_fn is not None else self.eviction_strategy.get_priority + + def _touch_full_eviction_key(self, node: UnifiedTreeNode) -> None: + """Call after any write to a Full eviction-key input + (last_access_time, hit_count, priority, Full session_ref). + One dict lookup for the common non-leaf case.""" + if node in self._full_device_live: + self.full_device_heap.refresh(node) + elif node in self._full_host_live: + self.full_host_heap.refresh(node) + def reset(self) -> None: """Rebuild the root, LRUs, sizes, evictable-leaf sets, and the empty match result.""" @@ -510,6 +714,23 @@ def reset(self) -> None: self.evictable_device_leaves: set[UnifiedTreeNode] = set() self.evictable_host_leaves: set[UnifiedTreeNode] = set() + # Persistent lazy heaps over the Full component's evictable leaves + # (see _LazyLeafHeap). Keys are resolved lazily on first use. + rebuild_each_walk = not envs.SGLANG_UNIFIED_RADIX_LAZY_EVICTION_HEAP.get() + self.full_device_heap = _LazyLeafHeap( + self.evictable_device_leaves, + self._full_eviction_key_fn, + rebuild_each_walk=rebuild_each_walk, + ) + self.full_host_heap = _LazyLeafHeap( + self.evictable_host_leaves, + self._full_eviction_key_fn, + rebuild_each_walk=rebuild_each_walk, + ) + # Direct references for the inlined "is this node live?" checks on the + # match/insert/lock hot paths (the heaps filter these dicts in place). + self._full_device_live = self.full_device_heap._live + self._full_host_live = self.full_host_heap._live self.host_lru_lists = { ct: UnifiedLRUList( ct, @@ -995,6 +1216,7 @@ def _match_post_processor( cur_time = get_and_increase_time_counter() while node_update: node_update.last_access_time = cur_time + self._touch_full_eviction_key(node_update) cur_time -= 0.00001 node_update = node_update.parent @@ -1063,6 +1285,7 @@ def collect_full_device_indices( def _touch_node(self, node: UnifiedTreeNode): node.last_access_time = get_and_increase_time_counter() + self._touch_full_eviction_key(node) if node != self.root_node: for comp in self.components: if comp.component_type == BASE_COMPONENT_TYPE: @@ -1076,6 +1299,7 @@ def _inc_hit_count_and_check(self, node: UnifiedTreeNode) -> bool: if self.is_write_back: return False node.hit_count += 1 + self._touch_full_eviction_key(node) if self.enable_external_cache_linker: return ( @@ -1261,6 +1485,7 @@ def _insert_walk_step(self, state: _InsertWalkState) -> None: if action is not None: step_actions.append(action) node.priority = max(node.priority, state.priority) + self._touch_full_eviction_key(node) if node.evicted: self._unevict_node_on_insert( @@ -1542,13 +1767,19 @@ def _update_evictable_leaf_sets(self, node: UnifiedTreeNode) -> None: """Update both device and host leaf sets for a node.""" if self._is_device_leaf(node): self.evictable_device_leaves.add(node) + self.full_device_heap.refresh(node) else: self.evictable_device_leaves.discard(node) + if node in self._full_device_live: + self.full_device_heap.forget(node) if self._is_host_leaf(node): self.evictable_host_leaves.add(node) + self.full_host_heap.refresh(node) else: self.evictable_host_leaves.discard(node) + if node in self._full_host_live: + self.full_host_heap.forget(node) def _update_duplicate_tracking(self, node: UnifiedTreeNode) -> None: """Register where duplicates are born (acks, split, unevict); @@ -1834,6 +2065,8 @@ def _release_all_component_layers( ) self.evictable_device_leaves.discard(node) self.evictable_host_leaves.discard(node) + self.full_device_heap.forget(node) + self.full_host_heap.forget(node) def _delete_unbacked_device_leaf( self, @@ -1983,6 +2216,7 @@ def _evict_host_leaf( ) tracker[comp.component_type] += hf self.evictable_host_leaves.discard(node) + self.full_host_heap.forget(node) self._remove_leaf_from_parent(node) self._iteratively_delete_tombstone_leaf(node, tracker, device_frees, host_frees) @@ -2229,6 +2463,7 @@ def _iteratively_delete_tombstone_ancestors( ) self.evictable_host_leaves.discard(cur) + self.full_host_heap.forget(cur) self._remove_leaf_from_parent(cur) parent = cur.parent self._update_evictable_leaf_sets(parent) @@ -3016,6 +3251,10 @@ def sanity_check( f"[Leaf] {len(overlap)} in both sets: {[n.id for n in list(overlap)[:5]]}" ) + # Persistent eviction heaps mirror the leaf sets with fresh keys. + self.full_device_heap.check_invariants(E, "D-heap") + self.full_host_heap.check_invariants(E, "H-heap") + if self.enable_session_radix_cache: for component in self.components: component.validate_session_state(all_node_set, E) diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py index c52cf12c5f7a..bdcb772d0385 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py @@ -398,6 +398,10 @@ def p99_us(self): idx = int(len(self.latencies_us) * 0.99) return sorted(self.latencies_us)[min(idx, len(self.latencies_us) - 1)] + @property + def mean_us(self): + return statistics.mean(self.latencies_us) if self.latencies_us else 0 + def report(self): tok = ( f"{self.tokens_per_sec:>12,.0f} tok/s" @@ -406,7 +410,8 @@ def report(self): ) return ( f" {self.name:<18s} | {tok} | {self.ops_per_sec:>10,.0f} ops/s | " - f"p50={self.p50_us:>8,.0f}us p99={self.p99_us:>8,.0f}us" + f"p50={self.p50_us:>8,.0f}us p99={self.p99_us:>8,.0f}us " + f"mean={self.mean_us:>8,.0f}us" ) @@ -551,6 +556,66 @@ def bench_evict( ) +def bench_evict_step( + num_seqs=5000, + chunk_len=256, + kv_size=500_000, + components=None, + verify=False, + page_size=1, + step_tokens=64, +): + """Decode-shaped eviction: pool full, then evict *step_tokens* and re-insert + one distinct sequence of the same length per step so the tree stays full. + + Isolates the per-call eviction overhead (heap maintenance over the + evictable-leaf set) instead of the prefill-shaped batch evictions of + ``bench_evict``. Prints the evictable-leaf count before and after. + """ + env = _make_env(num_seqs, chunk_len, kv_size, components, page_size) + inserted = _fill_no_evict(env) + step_tokens = max(step_tokens, page_size) + + def leaf_count(): + core = getattr(env.tree, "tree_core", None) + leaves = getattr(core, "evictable_device_leaves", None) + return len(leaves) if leaves is not None else -1 + + num_steps = min(1000, max(inserted // 5, 100)) + warmup = min(20, num_steps // 10) + # Slicing gen_random_sequences would repeat: every sequence shares the same + # chunk_len//4 root, so seq[:step_tokens] collapses to a handful of distinct + # payloads. Synthesize instead: a shared head keeps prefix reuse realistic, + # the per-step tail forces the new leaf the eviction walk has to manage. + head_len = max(1, min(step_tokens // 2, len(env.seqs[0]))) + head = list(env.seqs[0][:head_len]) + items = [ + (step_tokens, head + [1_000_000 + i] * (step_tokens - head_len)) + for i in range(num_steps + warmup) + ] + + def step(item): + n, seq = item + env.tree.evict(EvictParams(num_tokens=n, mamba_num=2)) + _insert_seq(env, seq) + + leaves_before = leaf_count() + result = bench_api( + "evict_step", + lambda: items, + step, + num_steps, + step_tokens, + warmup, + (lambda _: env.tree.sanity_check()) if verify else None, + ) + print( + f" evict_step | evictable leaves {leaves_before:,} -> {leaf_count():,} " + f"| {step_tokens} tokens/step" + ) + return result + + def bench_lock_unlock( num_seqs=5000, chunk_len=256, @@ -668,6 +733,7 @@ def bench_release( "insert": bench_insert, "match": bench_match_prefix, "evict": bench_evict, + "evict_step": bench_evict_step, "lock": bench_lock_unlock, "release": bench_release, } @@ -788,6 +854,9 @@ def test_bench_match_prefix(self): def test_bench_evict(self): self._run(bench_evict) + def test_bench_evict_step(self): + self._run(bench_evict_step) + def test_bench_lock_unlock(self): self._run(bench_lock_unlock) @@ -837,7 +906,7 @@ def _run_bench_cli(): "--benchmarks", nargs="+", default=["all"], - help="insert match evict lock release all", + help="insert match evict evict_step lock release all", ) args, _ = parser.parse_known_args() diff --git a/test/registered/unit/mem_cache/test_unified_radix_eviction_heap.py b/test/registered/unit/mem_cache/test_unified_radix_eviction_heap.py new file mode 100644 index 000000000000..d05dc7ff88d4 --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_radix_eviction_heap.py @@ -0,0 +1,639 @@ +"""Persistent lazy eviction heap for the UnifiedRadixCache Full component. + +Covers ``_LazyLeafHeap`` directly and proves that eviction order is identical +to the legacy per-call heap rebuild by replaying the same randomized op stream +on two caches: one running the new code and one whose Full component carries +the legacy ``_evict_device_*`` / ``drive_host_eviction`` bodies (installed +through ``component_registry_override``). +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=25, suite="base-a-test-cpu") + +import heapq +import random +import unittest +from array import array +from types import SimpleNamespace +from typing import Optional + +import torch +from numpy import float64 + +from sglang.srt.environ import envs +from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator +from sglang.srt.mem_cache.base_prefix_cache import ( + EvictParams, + InsertParams, + MatchPrefixParams, +) +from sglang.srt.mem_cache.cache_init_params import CacheInitParams +from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool +from sglang.srt.mem_cache.radix_cache import RadixKey +from sglang.srt.mem_cache.unified_cache.components import ( + ComponentType, +) +from sglang.srt.mem_cache.unified_cache.components import base as _tree_component +from sglang.srt.mem_cache.unified_cache.components.full import ( + FullComponent, +) +from sglang.srt.mem_cache.unified_cache.unified_tree_core import _LazyLeafHeap +from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache +from sglang.test.test_utils import CustomTestCase + +POLICIES = ("lru", "lfu", "fifo", "mru", "filo", "priority", "slru") + + +# --------------------------------------------------------------------------- +# Legacy oracle: the Full component as it was before the persistent heap. +# --------------------------------------------------------------------------- +class LegacyFullComponent(FullComponent): + """Verbatim pre-change eviction drivers (rebuild the heap every call).""" + + def _evict_device_start(self, request_cnt: int) -> None: + self._ensure_eviction_strategy() + self._evict_device_request_cnt = request_cnt + self._evict_device_last_node = None + self._evict_device_heap = [ + (self.session_ref_eviction_strategy(n), n) + for n in self.tree_core.evictable_device_leaves + ] + heapq.heapify(self._evict_device_heap) + + def _evict_device_next_node(self, tracker, device_frees, host_frees): + ct = self.component_type + lv = self._evict_device_last_node + if ( + lv is not None + and lv.parent is not None + and lv.parent in self.tree_core.evictable_device_leaves + ): + heapq.heappush( + self._evict_device_heap, + (self.session_ref_eviction_strategy(lv.parent), lv.parent), + ) + self._evict_device_last_node = None + while tracker[ct] < self._evict_device_request_cnt and self._evict_device_heap: + _, x = heapq.heappop(self._evict_device_heap) + if x not in self.tree_core.evictable_device_leaves: + continue + self._evict_device_last_node = x + return x.id + return None + + def _evict_device_end(self) -> None: + self._evict_device_heap = [] + self._evict_device_last_node = None + + def drive_host_eviction(self, num_tokens, tracker, device_frees, host_frees): + self._ensure_eviction_strategy() + heap = [ + (self.session_ref_eviction_strategy(n), n) + for n in self.tree_core.evictable_host_leaves + ] + heapq.heapify(heap) + ct = self.component_type + while tracker[ct] < num_tokens and heap: + _, x = heapq.heappop(heap) + if x not in self.tree_core.evictable_host_leaves: + continue + self.tree_core._evict_host_leaf(x, tracker, device_frees, host_frees) + if ( + x.parent is not None + and x.parent in self.tree_core.evictable_host_leaves + ): + heapq.heappush( + heap, + (self.session_ref_eviction_strategy(x.parent), x.parent), + ) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- +def make_cache( + *, + policy: str = "lru", + kv_size: int = 512, + enable_session: bool = False, + legacy: bool = False, + page_size: int = 1, +) -> UnifiedRadixCache: + dtype = torch.float16 + kv_pool = MHATokenToKVPool( + size=kv_size, + page_size=page_size, + dtype=dtype, + head_num=2, + head_dim=8, + layer_num=1, + device="cpu", + enable_memory_saver=False, + ) + allocator = TokenToKVPoolAllocator( + size=kv_size, dtype=dtype, device="cpu", kvcache=kv_pool, need_sort=False + ) + req_pool = ReqToTokenPool( + size=8, max_context_len=1024, device="cpu", enable_memory_saver=False + ) + # The heap lives in the Python TreeCore; pin it so the suite does not pick + # up the process-wide backend default. + with envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override("python"): + return UnifiedRadixCache( + params=CacheInitParams( + disable=False, + req_to_token_pool=req_pool, + token_to_kv_pool_allocator=allocator, + page_size=page_size, + eviction_policy=policy, + enable_session_radix_cache=enable_session, + tree_components=(ComponentType.FULL,), + component_registry_override=( + {ComponentType.FULL: LegacyFullComponent} if legacy else None + ), + ) + ) + + +def node_path(core, node) -> tuple: + """Root-to-node token path: identifies a node across two trees.""" + parts = [] + while node is not None and node is not core.root_node: + parts.append(tuple(node.key.token_ids)) + node = node.parent + return tuple(reversed(parts)) + + +def reset_time_counter() -> None: + _tree_component._LAST_ACCESS_TIME_COUNTER_FLOAT = float64(1.0) + + +def gen_sequences(rng: random.Random, n: int, page_size: int) -> list[list[int]]: + """Sequences with tree-like prefix sharing (chains + fan-out).""" + seqs = [[rng.randint(1, 50) for _ in range(page_size)]] + while len(seqs) < n: + parent = rng.choice(seqs) + tail = [rng.randint(1, 50) for _ in range(page_size * rng.randint(1, 6))] + seqs.append(parent + tail) + return seqs + + +def gen_ops(seed: int, num_ops: int, page_size: int, session: bool) -> list[tuple]: + rng = random.Random(seed) + seqs = gen_sequences(rng, 60, page_size) + ops = [] + for _ in range(num_ops): + r = rng.random() + if r < 0.40: + ops.append(("insert", rng.choice(seqs), rng.randint(0, 3))) + elif r < 0.60: + ops.append(("match", rng.choice(seqs))) + elif r < 0.70: + ops.append(("lock", rng.choice(seqs))) + elif r < 0.78: + ops.append(("unlock", rng.randint(0, 1 << 30))) + elif session and r < 0.84: + ops.append(("register", rng.choice(seqs), f"s{rng.randint(0, 3)}")) + elif session and r < 0.88: + ops.append(("release", f"s{rng.randint(0, 3)}")) + else: + ops.append(("evict", page_size * rng.randint(1, 6))) + return ops + + +class Replay: + """Drive one cache through an op stream, recording every eviction victim.""" + + def __init__(self, cache: UnifiedRadixCache, check_every: int = 25): + self.cache = cache + self.core = cache.tree_core + self.check_every = check_every + self.victims: list[list[tuple]] = [] + self.evicted_counts: list[int] = [] + self._current: Optional[list] = None + self._locked: list[tuple] = [] + orig = self.core.evict_device_leaf + core = self.core + + def recording_evict_device_leaf(node_id, is_write_back): + if self._current is not None: + self._current.append(node_path(core, core.node_by_id(node_id))) + return orig(node_id, is_write_back) + + self.core.evict_device_leaf = recording_evict_device_leaf + + def _alloc(self, n: int): + alloc = self.cache.token_to_kv_pool_allocator + v = alloc.alloc(n) + if v is None: + self.cache.evict(EvictParams(num_tokens=n * 2)) + v = alloc.alloc(n) + return v + + def run(self, ops: list[tuple]) -> None: + for i, op in enumerate(ops): + kind = op[0] + if kind == "insert": + seq, prio = op[1], op[2] + v = self._alloc(len(seq)) + if v is not None: + self.cache.insert( + InsertParams( + key=RadixKey(array("q", seq)), + value=v.to(torch.int64), + priority=prio, + ) + ) + elif kind == "match": + self.cache.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", op[1]))) + ) + elif kind == "lock": + res = self.cache.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", op[1]))) + ) + node_id = res.last_device_node + if node_id != self.cache.root_node_handle(): + lr = self.cache.inc_lock_ref(node_id).to_dec_params() + self._locked.append((node_id, lr)) + elif kind == "unlock": + if self._locked: + node_id, lr = self._locked.pop(op[1] % len(self._locked)) + self.cache.dec_lock_ref(node_id, lr) + elif kind == "register": + seq, sid = op[1], op[2] + res = self.cache.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", seq))) + ) + if res.last_device_node != self.cache.root_node_handle(): + self.cache.session_refs.register_session_ref( + SimpleNamespace( + session_id=sid, + session_generation=self.cache.ensure_session_generation( + sid + ), + session=None, + last_node=res.last_device_node, + origin_input_ids=array("q", seq), + output_ids=array("q"), + extra_key=None, + ), + leaf=res.last_device_node, + ) + elif kind == "release": + self.cache.release_radix_session(op[1]) + elif kind == "evict": + self._current = [] + res = self.cache.evict(EvictParams(num_tokens=op[1])) + self.victims.append(self._current) + self.evicted_counts.append(res.num_tokens_evicted) + self._current = None + if i % self.check_every == 0: + self.cache.sanity_check() + # release everything so the final state is comparable + while self._locked: + node_id, lr = self._locked.pop() + self.cache.dec_lock_ref(node_id, lr) + self.cache.sanity_check() + + def leaf_paths(self) -> set: + return {node_path(self.core, n) for n in self.core.evictable_device_leaves} + + +def replay_pair( + policy: str, seed: int, session: bool, page_size: int = 1, kill_switch=False +): + ops = gen_ops(seed, 320, page_size, session) + reset_time_counter() + # Pin the lazy side explicitly: the env var is read at UnifiedTreeCore + # construction, so an exported kill switch would otherwise flip both sides. + with envs.SGLANG_UNIFIED_RADIX_LAZY_EVICTION_HEAP.override(True): + new = Replay( + make_cache(policy=policy, enable_session=session, page_size=page_size) + ) + new.run(ops) + reset_time_counter() + if kill_switch: + with envs.SGLANG_UNIFIED_RADIX_LAZY_EVICTION_HEAP.override(False): + ref = Replay( + make_cache(policy=policy, enable_session=session, page_size=page_size) + ) + else: + ref = Replay( + make_cache( + policy=policy, enable_session=session, page_size=page_size, legacy=True + ) + ) + ref.run(ops) + return new, ref + + +# --------------------------------------------------------------------------- +# _LazyLeafHeap unit tests +# --------------------------------------------------------------------------- +class _Node: + __slots__ = ("id", "key", "parent") + + def __init__(self, id, key): + self.id = id + self.key = key + self.parent = None + + def __lt__(self, other): + return self.id < other.id + + def __repr__(self): + return f"N{self.id}" + + +class TestLazyLeafHeapUnit(CustomTestCase): + def _heap(self, nodes, **kw): + members = set(nodes) + heap = _LazyLeafHeap(members, lambda: lambda n: n.key, **kw) + for n in nodes: + heap.refresh(n) + return members, heap + + def _drain(self, heap): + heap.begin_walk() + out = [] + while True: + n = heap.pop_next() + if n is None: + break + out.append(n) + heap.end_walk() + return out + + def test_pop_order_and_key_updates(self): + nodes = [_Node(i, k) for i, k in enumerate([5, 3, 9, 1])] + members, heap = self._heap(nodes) + errors = [] + heap.check_invariants(errors.append, "t") + self.assertEqual(errors, []) + nodes[2].key = 0 # 9 -> 0: must come first after a touch + heap.touch(nodes[2]) + self.assertEqual(self._drain(heap), [nodes[2], nodes[3], nodes[1], nodes[0]]) + # a full walk yields every member exactly once, then they are all live again + errors = [] + heap.check_invariants(errors.append, "t") + self.assertEqual(errors, []) + + def test_touch_ignores_non_members_and_forget_removes(self): + nodes = [_Node(i, i) for i in range(4)] + members, heap = self._heap(nodes) + outsider = _Node(99, -1) + heap.touch(outsider) + heap.refresh(outsider) + self.assertEqual(len(heap), 4) + members.discard(nodes[0]) + heap.forget(nodes[0]) + self.assertEqual(self._drain(heap), nodes[1:]) + + def test_walk_snapshot_entrants_and_promote(self): + nodes = [_Node(i, i) for i in range(3)] + members, heap = self._heap(nodes) + heap.begin_walk() + first = heap.pop_next() + self.assertIs(first, nodes[0]) + entrant = _Node(10, -5) # smallest key, enters mid-walk + members.add(entrant) + heap.refresh(entrant) + self.assertIs(heap.pop_next(), nodes[1]) # invisible without promote + heap.promote(entrant) + self.assertIs(heap.pop_next(), entrant) # visible right after promote + self.assertIs(heap.pop_next(), nodes[2]) + self.assertIsNone(heap.pop_next()) + heap.end_walk() + # every member (including the yielded ones) is live again + self.assertEqual(set(heap._live), members) + errors = [] + heap.check_invariants(errors.append, "t") + self.assertEqual(errors, []) + + def test_mid_walk_touch_is_deferred_to_end_walk(self): + nodes = [_Node(i, i) for i in range(3)] + members, heap = self._heap(nodes) + heap.begin_walk() + nodes[2].key = -100 + heap.touch(nodes[2]) # key frozen for this walk + self.assertIs(heap.pop_next(), nodes[0]) + heap.end_walk() + # after the walk the fresh key is honoured; the yielded node is live again + self.assertEqual(self._drain(heap), [nodes[2], nodes[0], nodes[1]]) + + def test_nested_walk_is_rejected(self): + members, heap = self._heap([_Node(0, 0)]) + heap.begin_walk() + with self.assertRaises(AssertionError): + heap.begin_walk() + heap.end_walk() + + def test_compaction_bound(self): + nodes = [_Node(i, i) for i in range(8)] + members, heap = self._heap(nodes) + for t in range(5000): + nodes[t % 8].key = 1000 + t + heap.touch(nodes[t % 8]) + self.assertLessEqual(len(heap._heap), 2 * len(heap._live) + 64) + errors = [] + heap.check_invariants(errors.append, "t") + self.assertEqual(errors, []) + self.assertEqual(self._drain(heap), sorted(nodes, key=lambda n: n.key)) + + def test_rebuild_each_walk_matches(self): + nodes = [_Node(i, k) for i, k in enumerate([4, 2, 8, 6])] + _, lazy = self._heap(nodes) + _, eager = self._heap(nodes, rebuild_each_walk=True) + nodes[0].key = 1 # key change without any hook: only the rebuild sees it + self.assertEqual(self._drain(eager)[0], nodes[0]) + self.assertEqual(self._drain(lazy)[0], nodes[1]) + + def test_check_invariants_reports_divergence(self): + nodes = [_Node(i, i) for i in range(3)] + members, heap = self._heap(nodes) + members.discard(nodes[1]) # membership removed without forget() + errors = [] + heap.check_invariants(errors.append, "t") + self.assertTrue(any("live but not member" in e for e in errors)) + members.add(nodes[1]) + nodes[2].key = 42 # key changed without touch() + errors = [] + heap.check_invariants(errors.append, "t") + self.assertTrue(any("out of date" in e for e in errors)) + + +# --------------------------------------------------------------------------- +# Order parity against the legacy per-call rebuild +# --------------------------------------------------------------------------- +class TestEvictionOrderParity(CustomTestCase): + def _assert_parity(self, new: Replay, ref: Replay, label: str): + self.assertEqual(len(new.victims), len(ref.victims), label) + for i, (a, b) in enumerate(zip(new.victims, ref.victims)): + self.assertEqual(a, b, f"{label}: victims differ at evict call {i}") + self.assertEqual(new.evicted_counts, ref.evicted_counts, label) + self.assertEqual(new.leaf_paths(), ref.leaf_paths(), label) + self.assertGreater(sum(len(v) for v in new.victims), 0, label) + + def test_parity_all_policies(self): + for policy in POLICIES: + for session in (False, True): + for seed in range(12): + label = f"policy={policy} session={session} seed={seed}" + with self.subTest(label): + new, ref = replay_pair(policy, seed, session) + self._assert_parity(new, ref, label) + + def test_parity_page_size_4(self): + for policy in ("lru", "lfu", "mru"): + for seed in range(2): + label = f"page_size=4 policy={policy} seed={seed}" + with self.subTest(label): + new, ref = replay_pair(policy, seed, False, page_size=4) + self._assert_parity(new, ref, label) + + def test_kill_switch_parity(self): + for policy in ("lru", "slru"): + label = f"kill-switch policy={policy}" + with self.subTest(label): + new, ref = replay_pair(policy, 7, True, kill_switch=True) + self._assert_parity(new, ref, label) + self.assertTrue(ref.core.full_device_heap.rebuild_each_walk) + self.assertFalse(new.core.full_device_heap.rebuild_each_walk) + + +# --------------------------------------------------------------------------- +# Targeted behaviours on a real cache +# --------------------------------------------------------------------------- +def _insert(cache, tokens): + v = cache.token_to_kv_pool_allocator.alloc(len(tokens)) + cache.insert( + InsertParams(key=RadixKey(array("q", tokens)), value=v.to(torch.int64)) + ) + + +def _match_len(cache, tokens) -> int: + return len( + cache.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens))) + ).device_indices + ) + + +class TestHeapOnRealCache(CustomTestCase): + def test_mru_touched_leaf_is_evicted_first(self): + cache = make_cache(policy="mru") + _insert(cache, [1, 2, 3]) + _insert(cache, [7, 8, 9]) + _match_len(cache, [7, 8, 9]) # most recently used -> first victim under MRU + cache.evict(EvictParams(num_tokens=3)) + self.assertEqual(_match_len(cache, [7, 8, 9]), 0) + self.assertEqual(_match_len(cache, [1, 2, 3]), 3) + cache.sanity_check() + + def test_lru_touched_leaf_is_protected(self): + cache = make_cache(policy="lru") + _insert(cache, [1, 2, 3]) + _insert(cache, [7, 8, 9]) + _match_len(cache, [1, 2, 3]) # refresh -> [7,8,9] is the LRU victim + cache.evict(EvictParams(num_tokens=3)) + self.assertEqual(_match_len(cache, [7, 8, 9]), 0) + self.assertEqual(_match_len(cache, [1, 2, 3]), 3) + cache.sanity_check() + + def test_session_release_returns_leaf_to_unreferenced_band(self): + cache = make_cache(policy="lru", enable_session=True) + _insert(cache, [1, 2, 3, 4]) + _insert(cache, [7, 8, 9]) + res = cache.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3, 4]))) + ) + cache.session_refs.register_session_ref( + SimpleNamespace( + session_id="s1", + session_generation=cache.ensure_session_generation("s1"), + session=None, + last_node=res.last_device_node, + origin_input_ids=array("q", [1, 2, 3, 4]), + output_ids=array("q"), + extra_key=None, + ), + leaf=res.last_device_node, + ) + cache.evict(EvictParams(num_tokens=3)) + self.assertEqual(_match_len(cache, [7, 8, 9]), 0) # unreferenced goes first + self.assertEqual(_match_len(cache, [1, 2, 3, 4]), 4) + cache.release_radix_session("s1") + cache.sanity_check() + cache.evict(EvictParams(num_tokens=4)) + self.assertEqual( + _match_len(cache, [1, 2, 3, 4]), 0 + ) # released leaf is evictable again + cache.sanity_check() + + def test_parent_promotion_within_one_call(self): + cache = make_cache(policy="lru") + _insert(cache, [1, 2]) + _insert(cache, [1, 2, 3, 4]) + # One call for 4 tokens must evict the leaf [3,4] and then its parent [1,2]. + res = cache.evict(EvictParams(num_tokens=4)) + self.assertEqual(res.num_tokens_evicted, 4) + self.assertEqual(_match_len(cache, [1, 2]), 0) + cache.sanity_check() + + def test_reset_empties_heaps(self): + cache = make_cache(policy="lru") + _insert(cache, [1, 2, 3]) + self.assertEqual(len(cache.tree_core.full_device_heap), 1) + cache.reset() + self.assertEqual(len(cache.tree_core.full_device_heap), 0) + self.assertEqual(len(cache.tree_core.full_host_heap), 0) + cache.sanity_check() + + def test_exception_mid_walk_leaves_heap_consistent(self): + cache = make_cache(policy="lru") + _insert(cache, [1, 2, 3]) + _insert(cache, [7, 8, 9]) + core = cache.tree_core + orig = core.evict_device_leaf + + def boom(node_id, is_write_back): + raise RuntimeError("injected") + + core.evict_device_leaf = boom + with self.assertRaises(RuntimeError): + cache.evict(EvictParams(num_tokens=3)) + core.evict_device_leaf = orig + self.assertFalse(core.full_device_heap.walking) + cache.sanity_check() + cache.evict(EvictParams(num_tokens=6)) + self.assertEqual(_match_len(cache, [1, 2, 3]) + _match_len(cache, [7, 8, 9]), 0) + cache.sanity_check() + + def test_eviction_walk_keeps_heap_compact(self): + """Regression: a walk must not leave one stale entry per evicted leaf.""" + cache = make_cache(policy="mru") + seqs = [[1000 + i * 3 + j for j in range(3)] for i in range(30)] + for s in seqs: + _insert(cache, s) + for _ in range(3): + for s in seqs: + _match_len(cache, s) + cache.evict(EvictParams(num_tokens=45)) + heap = cache.tree_core.full_device_heap + self.assertLessEqual(len(heap._heap), 2 * len(heap._live) + 64) + cache.sanity_check() + + def test_many_matches_keep_heap_compact(self): + cache = make_cache(policy="lru") + _insert(cache, [1, 2, 3]) + for _ in range(5000): + _match_len(cache, [1, 2, 3]) + heap = cache.tree_core.full_device_heap + self.assertLessEqual(len(heap._heap), 2 * len(heap._live) + 64) + cache.sanity_check() + + +if __name__ == "__main__": + unittest.main()