diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index de64ad8ab..d90cbadd3 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -548,7 +548,7 @@ class HiMambaRadixCache(MambaRadixCache): self._discard_from_leaf_sets(node) parent = node.parent - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -584,7 +584,7 @@ class HiMambaRadixCache(MambaRadixCache): self._discard_from_leaf_sets(node) parent = node.parent - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -598,7 +598,7 @@ class HiMambaRadixCache(MambaRadixCache): assert node.mamba_host_value is None, f"has mamba host value, {node.id=}" assert len(node.children) == 0, f"leaf node has children, {node.id=}" parent = node.parent - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -824,7 +824,7 @@ class HiMambaRadixCache(MambaRadixCache): if len(key) == 0: return 0, True - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children.keys(): @@ -836,7 +836,7 @@ class HiMambaRadixCache(MambaRadixCache): if node.mamba_value is not None: self.mamba_lru_list.reset_node_mru(node) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) if prefix_len < len(node.key): new_node = self._split_node(node.key, node, prefix_len) @@ -855,7 +855,7 @@ class HiMambaRadixCache(MambaRadixCache): value = value[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) mamba_value_exist = False if len(key): @@ -884,7 +884,7 @@ class HiMambaRadixCache(MambaRadixCache): value: torch.Tensor, mamba_value: torch.Tensor, ) -> TreeNode: - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) new_node = TreeNode() new_node.parent = parent new_node.key = key @@ -924,7 +924,7 @@ class HiMambaRadixCache(MambaRadixCache): ) -> Tuple[List[torch.Tensor], TreeNode, int]: """Walk tree to find best_last_node (mamba boundary).""" node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value: List[torch.Tensor] = [] best_value_len = 0 @@ -940,7 +940,7 @@ class HiMambaRadixCache(MambaRadixCache): best_value_len = len(value) best_last_node = node - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): new_node = self._split_node(child.key, child, prefix_len) if not new_node.evicted: @@ -953,7 +953,7 @@ class HiMambaRadixCache(MambaRadixCache): node = child key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) if node.mamba_value is not None or node.mamba_backuped: best_value_len = len(value) @@ -1074,7 +1074,7 @@ class HiMambaRadixCache(MambaRadixCache): self.evictable_full_host_leaves.discard(child) new_node = TreeNode() - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.value = None new_node.mamba_value = None @@ -1095,7 +1095,7 @@ class HiMambaRadixCache(MambaRadixCache): self.mamba_lru_list.remove_node(child) child.parent = new_node child.key = child.key[split_len:] - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node if child.mamba_value is not None: self.mamba_lru_list.insert_mru(child) @@ -1836,7 +1836,7 @@ class HiMambaRadixCache(MambaRadixCache): if len(key) == 0: return 0 - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) matched_length = 0 while len(key) > 0 and child_key in node.children.keys(): @@ -1844,7 +1844,7 @@ class HiMambaRadixCache(MambaRadixCache): node.last_access_time = get_last_access_time() if node != self.root_node and node.mamba_value is not None: self.mamba_lru_list.reset_node_mru(node) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) key = key[prefix_len:] host_value = host_value[prefix_len:] @@ -1856,7 +1856,7 @@ class HiMambaRadixCache(MambaRadixCache): node = new_node if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) leaf_node: Optional[TreeNode] = None if len(key): diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index c8e869cc4..6b2c1fe21 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -781,11 +781,7 @@ class HiRadixCache(RadixCache): return self.evictable_size_ def _to_radix_key(self, token_ids: List[int]) -> RadixKey: - """Convert raw token_ids to a RadixKey for tree walking. - - Must use list (not tuple) to match scheduler's RadixKey format, - since _key_match_paged compares slices directly and list != tuple. - """ + """Convert raw token_ids to a RadixKey; must be list (not tuple) for paged match.""" return RadixKey(token_ids=list(token_ids)) def inc_lock_ref(self, node: TreeNode) -> IncLockRefResult: @@ -938,7 +934,7 @@ class HiRadixCache(RadixCache): self._record_remove_event(x, medium=StorageMedium.CPU) num_evicted += self.cache_controller.evict_host(x.host_value) - key = self.get_child_key_fn(x.key) + key = x.key.child_key(self.page_size) v = x.parent.children.pop(key, None) assert v == x, f"parent does not have child key, {key}" if x in self.evictable_host_leaves: @@ -1312,13 +1308,13 @@ class HiRadixCache(RadixCache): if len(key) == 0: return 0 - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) matched_length = 0 while len(key) > 0 and child_key in node.children.keys(): node = node.children[child_key] node.last_access_time = time.monotonic() - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) key = key[prefix_len:] host_value = host_value[prefix_len:] hash_value = hash_value[prefix_len // self.page_size :] @@ -1329,7 +1325,7 @@ class HiRadixCache(RadixCache): node = new_node if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) if len(key): new_node = TreeNode(priority=node.priority) @@ -1350,13 +1346,13 @@ class HiRadixCache(RadixCache): def _match_prefix_helper(self, node: TreeNode, key: RadixKey): node.last_access_time = time.monotonic() - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value = [] while len(key) > 0 and child_key in node.children.keys(): child = node.children[child_key] child.last_access_time = time.monotonic() - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): new_node = self._split_node(child.key, child, prefix_len) if not new_node.evicted: @@ -1370,14 +1366,14 @@ class HiRadixCache(RadixCache): key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) return value, node def _split_node(self, key: RadixKey, child: TreeNode, split_len: int): # child node split into new_node -> child new_node = TreeNode(priority=child.priority) - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.lock_ref = child.lock_ref new_node.key = child.key[:split_len] @@ -1398,7 +1394,7 @@ class HiRadixCache(RadixCache): ) child.parent = new_node child.key = child.key[split_len:] - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node return new_node @@ -1420,14 +1416,14 @@ class HiRadixCache(RadixCache): return InsertResult(prefix_len=0) node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children.keys(): node = node.children[child_key] node.last_access_time = time.monotonic() node.priority = max(node.priority, priority) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) if prefix_len == len(node.key): if node.evicted: @@ -1463,7 +1459,7 @@ class HiRadixCache(RadixCache): value = value[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) if len(key): new_node = TreeNode(priority=priority) diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index d07702cf1..d201093f2 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -21,7 +21,7 @@ The radix tree data structure for managing the hybrid (full and Mamba) KV cache. import heapq from collections import defaultdict -from functools import lru_cache, partial +from functools import lru_cache from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -46,12 +46,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchResult, ) from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool -from sglang.srt.mem_cache.radix_cache import ( - RadixKey, - _key_match_page_size1, - _key_match_paged, - get_child_key, -) +from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.server_args import get_global_server_args if TYPE_CHECKING: @@ -445,12 +440,6 @@ class MambaRadixCache(BasePrefixCache): if params.enable_metrics: self.init_metrics_collector() - if self.page_size == 1: - self.key_match_fn = _key_match_page_size1 - self.get_child_key_fn = get_child_key - else: - self.key_match_fn = partial(_key_match_paged, page_size=self.page_size) - self.get_child_key_fn = partial(get_child_key, page_size=self.page_size) self.reset() ##### Public API ##### @@ -963,7 +952,7 @@ class MambaRadixCache(BasePrefixCache): node is greater than or equal to the sliding window size. """ node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value: List[torch.Tensor] = [] best_value_len = 0 @@ -975,7 +964,7 @@ class MambaRadixCache(BasePrefixCache): best_value_len = len(value) best_last_node = node - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): new_node = self._split_node(child.key, child, prefix_len) value.append(new_node.value) @@ -987,7 +976,7 @@ class MambaRadixCache(BasePrefixCache): key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) # handle best_value_len and best_last_node, for the case that last node is fully matched if node.mamba_value is not None: best_value_len = len(value) @@ -1081,7 +1070,7 @@ class MambaRadixCache(BasePrefixCache): def _split_node(self, key: RadixKey, child: TreeNode, split_len: int) -> TreeNode: # new_node -> child new_node = TreeNode() - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.mamba_value = None # mamba cache can not be split new_node.full_lock_ref = child.full_lock_ref @@ -1098,7 +1087,7 @@ class MambaRadixCache(BasePrefixCache): child.parent = new_node child.key = child.key[split_len:] child.value = child.value[split_len:].clone() - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node # insert the new node and child into the lru lists, insert # parent first so that parent is after child in the lru list @@ -1128,7 +1117,7 @@ class MambaRadixCache(BasePrefixCache): if len(key) == 0: return 0, True - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children.keys(): @@ -1137,7 +1126,7 @@ class MambaRadixCache(BasePrefixCache): self.full_lru_list.reset_node_mru(node) if node.mamba_value is not None: self.mamba_lru_list.reset_node_mru(node) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) if prev_prefix_len < total_prefix_length + prefix_len: start = max(0, prev_prefix_len - total_prefix_length) @@ -1152,7 +1141,7 @@ class MambaRadixCache(BasePrefixCache): node = new_node if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) mamba_value_exist = False if len(key): @@ -1208,7 +1197,7 @@ class MambaRadixCache(BasePrefixCache): node.mamba_value is not None ), f"Invariant violated: leaf node is a tombstone, {node.id=}" assert len(node.children) == 0, f"leaf node has children, {node.id=}" - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -1225,7 +1214,7 @@ class MambaRadixCache(BasePrefixCache): node.mamba_value is None ), f"Deleting a unexpected non-tombstone leaf node, {node.id=}" assert len(node.children) == 0, f"leaf node has children, {node.id=}" - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -1270,9 +1259,9 @@ class MambaRadixCache(BasePrefixCache): for key, child in current_node.children.items(): stack.append((child, current_indent + 2)) - assert key == self.get_child_key_fn( - child.key - ), f"{key=}, {self.get_child_key_fn(child.key)=}" + assert key == child.key.child_key( + self.page_size + ), f"{key=}, {child.key.child_key(self.page_size)=}" def _total_size_helper(self) -> Tuple[int, int]: total_size = 0 diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index 3e01bca91..06fec7547 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -21,12 +21,13 @@ limitations under the License. The radix tree data structure for managing the KV cache. """ +import hashlib import heapq import logging import sys import time from collections import defaultdict -from functools import lru_cache, partial +from functools import lru_cache from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union import torch @@ -141,6 +142,71 @@ class RadixKey: value = value[: len(self)] return self, value + def _check_compatible(self, other: "RadixKey") -> None: + if self.extra_key != other.extra_key: + raise ValueError( + f"RadixKey operations require matching extra_key, but got " + f"{self.extra_key=} != {other.extra_key=}" + ) + + def match(self, other: "RadixKey", page_size: int = 1) -> int: + """Logical-unit prefix length shared with ``other``. Result is rounded down to ``page_size``.""" + self._check_compatible(other) + t0, t1 = self.token_ids, other.token_ids + + if self.is_bigram: + # Walk raw tokens; L matching tokens imply L-1 matching bigrams. + i = 0 + for a, b in zip(t0, t1): + if a != b: + break + i += 1 + matched = max(0, min(i - 1, len(self), len(other))) + return (matched // page_size) * page_size if page_size > 1 else matched + + if page_size == 1: + i = 0 + for a, b in zip(t0, t1): + if a != b: + break + i += 1 + return i + + min_len = min(len(self), len(other)) + i = 0 + while i < min_len: + if t0[i : i + page_size] != t1[i : i + page_size]: + break + i += page_size + return i + + def child_key(self, page_size: int = 1): + """Hashable dict-key for the first ``page_size`` logical units, namespaced by ``extra_key``.""" + t = self.token_ids + if self.is_bigram: + if page_size == 1: + plain = (t[0], t[1]) + else: + plain = tuple((t[j], t[j + 1]) for j in range(page_size)) + else: + plain = t[0] if page_size == 1 else tuple(t[:page_size]) + return plain if self.extra_key is None else (self.extra_key, plain) + + def hash_page(self, start: int, end: int, prior_hash: Optional[str] = None) -> str: + """SHA256 for logical units [start, end); bigram mode feeds overlapping (t_i, t_{i+1}) byte pairs.""" + hasher = hashlib.sha256() + if prior_hash: + hasher.update(bytes.fromhex(prior_hash)) + t = self.token_ids + if self.is_bigram: + for j in range(start, end): + hasher.update(t[j].to_bytes(4, byteorder="little", signed=False)) + hasher.update(t[j + 1].to_bytes(4, byteorder="little", signed=False)) + else: + for j in range(start, end): + hasher.update(t[j].to_bytes(4, byteorder="little", signed=False)) + return hasher.hexdigest() + class TreeNode: @@ -205,82 +271,8 @@ class TreeNode: return self.last_access_time < other.last_access_time -def _check_extra_key(key0: RadixKey, key1: RadixKey): - if key0.extra_key != key1.extra_key: - raise ValueError( - f"_key_match should be run on the same extra key, but got key0.extra_key={key0.extra_key} != key1.extra_key={key1.extra_key}" - ) - - -def _key_match_page_size1(key0: RadixKey, key1: RadixKey): - _check_extra_key(key0, key1) - # In bigram mode we compare raw tokens position-by-position; matching L - # consecutive tokens implies L-1 matching bigrams. In plain mode, matching - # tokens == matching units directly. - t0 = key0.token_ids - t1 = key1.token_ids - i = 0 - for a, b in zip(t0, t1): - if a != b: - break - i += 1 - if key0.is_bigram: - # Clamp by logical bigram length of each side (guards short tails). - return max(0, min(i - 1, len(key0), len(key1))) - return i - - -def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int): - _check_extra_key(key0, key1) - if key0.is_bigram: - # Walk raw tokens, convert to bigram count, then round to page boundary. - t0 = key0.token_ids - t1 = key1.token_ids - i = 0 - for a, b in zip(t0, t1): - if a != b: - break - i += 1 - bigram_matched = max(0, i - 1) - bigram_matched = min(bigram_matched, len(key0), len(key1)) - return (bigram_matched // page_size) * page_size - - min_len = min(len(key0), len(key1)) - i = 0 - while i < min_len: - if key0.token_ids[i : i + page_size] != key1.token_ids[i : i + page_size]: - break - i += page_size - return i - - -def get_child_key(key: RadixKey, page_size: int = 1): - if key.is_bigram: - t = key.token_ids - if page_size == 1: - # first bigram -> (tokens[0], tokens[1]) - plain_key = (t[0], t[1]) - else: - # first page_size bigrams spanning tokens[0 : page_size + 1] - plain_key = tuple((t[j], t[j + 1]) for j in range(page_size)) - else: - if page_size == 1: - plain_key = key.token_ids[0] - else: - plain_key = tuple(key.token_ids[:page_size]) - if key.extra_key is None: - return plain_key - else: - return (key.extra_key, plain_key) - - def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]: - """Compute SHA256-based hash values for position-aware identification. - - In bigram mode, each page logically covers `page_size` bigrams over - `page_size + 1` raw tokens; we feed overlapping (t_i, t_{i+1}) byte pairs - to the hasher so the output matches the pre-optimization tuple-based hash. - """ + """Compute SHA256-based hash values for position-aware identification.""" hash_values = [] parent_hash = None @@ -288,45 +280,17 @@ def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]: if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0: parent_hash = node.parent.hash_value[-1] - raw = node.key.token_ids - is_bigram = node.key.is_bigram logical_len = len(node.key) - for start in range(0, logical_len, page_size): end = min(start + page_size, logical_len) if end <= start: continue - hash_val = _hash_page(raw, start, end, is_bigram, parent_hash) + hash_val = node.key.hash_page(start, end, parent_hash) hash_values.append(hash_val) parent_hash = hash_val - return hash_values -def _hash_page( - raw_tokens: List[int], - start: int, - end: int, - is_bigram: bool, - prior_hash: Optional[str], -) -> str: - import hashlib - - hasher = hashlib.sha256() - if prior_hash: - hasher.update(bytes.fromhex(prior_hash)) - if is_bigram: - for j in range(start, end): - hasher.update(raw_tokens[j].to_bytes(4, byteorder="little", signed=False)) - hasher.update( - raw_tokens[j + 1].to_bytes(4, byteorder="little", signed=False) - ) - else: - for j in range(start, end): - hasher.update(raw_tokens[j].to_bytes(4, byteorder="little", signed=False)) - return hasher.hexdigest() - - def split_node_hash_value( child_hash_value: Optional[List[str]], split_len: int, page_size: int ) -> tuple[Optional[List[str]], Optional[List[str]]]: @@ -375,13 +339,6 @@ class RadixCache(BasePrefixCache): else: self.device = torch.device("cpu") - if self.page_size == 1: - self.key_match_fn = _key_match_page_size1 - self.get_child_key_fn = get_child_key - else: - self.key_match_fn = partial(_key_match_paged, page_size=self.page_size) - self.get_child_key_fn = partial(get_child_key, page_size=self.page_size) - if self.eviction_policy == "lru": self.eviction_strategy: EvictionStrategy = LRUStrategy() elif self.eviction_policy == "lfu": @@ -737,13 +694,13 @@ class RadixCache(BasePrefixCache): access_time = time.monotonic() node.last_access_time = access_time - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value = [] while len(key) > 0 and child_key in node.children.keys(): child = node.children[child_key] child.last_access_time = access_time - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): new_node = self._split_node(child.key, child, prefix_len) value.append(new_node.value) @@ -755,7 +712,7 @@ class RadixCache(BasePrefixCache): key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) return value, node @@ -764,7 +721,7 @@ class RadixCache(BasePrefixCache): # New node inherits child's priority (represents shared prefix) new_node = TreeNode(priority=child.priority) new_node.hit_count = child.hit_count - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.lock_ref = child.lock_ref new_node.key = child.key[:split_len] @@ -772,7 +729,7 @@ class RadixCache(BasePrefixCache): child.parent = new_node child.key = child.key[split_len:] child.value = child.value[split_len:].clone() - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node # Split hash_value if it was already computed, otherwise leave as None new_node.hash_value, child.hash_value = split_node_hash_value( @@ -807,13 +764,13 @@ class RadixCache(BasePrefixCache): if len(key) == 0: return 0 - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children.keys(): node = node.children[child_key] node.last_access_time = access_time - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) total_prefix_length += prefix_len key = key[prefix_len:] value = value[prefix_len:] @@ -827,7 +784,7 @@ class RadixCache(BasePrefixCache): node.priority = max(node.priority, priority) self._inc_hit_count(node, chunked) if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) if len(key): new_node = TreeNode(priority=priority) @@ -857,12 +814,12 @@ class RadixCache(BasePrefixCache): for key, child in current_node.children.items(): stack.append((child, current_indent + 2)) - assert key == self.get_child_key_fn( - child.key - ), f"{key=}, {self.get_child_key_fn(child.key)=}" + assert key == child.key.child_key( + self.page_size + ), f"{key=}, {child.key.child_key(self.page_size)=}" def _delete_leaf(self, node): - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" diff --git a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py index 9a82aa31f..da0d4cf17 100644 --- a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py @@ -194,7 +194,7 @@ class LMCRadixCache(RadixCache): new_node.key = key[start:end] new_node.value = token_slots[:fetched] new_node.parent = last_node - last_node.children[self.get_child_key_fn(new_node.key)] = new_node + last_node.children[new_node.key.child_key(self.page_size)] = new_node last_node = new_node value = torch.cat([value, token_slots[:fetched]]) diff --git a/python/sglang/srt/mem_cache/swa_radix_cache.py b/python/sglang/srt/mem_cache/swa_radix_cache.py index b5450015e..c650588b6 100644 --- a/python/sglang/srt/mem_cache/swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/swa_radix_cache.py @@ -22,7 +22,6 @@ The radix tree data structure for managing the hybrid (full and SWA) KV cache. import heapq import time from collections import defaultdict -from functools import partial from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -41,12 +40,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchResult, ) from sglang.srt.mem_cache.cache_init_params import CacheInitParams -from sglang.srt.mem_cache.radix_cache import ( - RadixKey, - _key_match_page_size1, - _key_match_paged, - get_child_key, -) +from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.utils import convert_to_bigram_key @@ -353,13 +347,6 @@ class SWARadixCache(BasePrefixCache): else: self.device = torch.device("cpu") - if self.page_size == 1: - self.key_match_fn = _key_match_page_size1 - self.get_child_key_fn = get_child_key - else: - self.key_match_fn = partial(_key_match_paged, page_size=self.page_size) - self.get_child_key_fn = partial(get_child_key, page_size=self.page_size) - if self.is_eagle: self.key_convert_fn = convert_to_bigram_key else: @@ -795,7 +782,7 @@ class SWARadixCache(BasePrefixCache): node is greater than or equal to the sliding window size. """ node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value = [] # for path connected to root without tombstone, always match, so set to inf @@ -813,7 +800,7 @@ class SWARadixCache(BasePrefixCache): # reset match_len_since_tombstone if we hit a tombstone node match_len_since_tombstone = 0 - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): new_node = self._split_node(child.key, child, prefix_len) value.append(new_node.value) @@ -829,7 +816,7 @@ class SWARadixCache(BasePrefixCache): key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) # handle best_value_len and best_last_node, for the case that last node is fully matched if match_len_since_tombstone >= self.sliding_window_size: @@ -887,7 +874,7 @@ class SWARadixCache(BasePrefixCache): def _split_node(self, key: RadixKey, child: TreeNode, split_len: int) -> TreeNode: # new_node -> child new_node = TreeNode() - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.swa_tombstone = child.swa_tombstone new_node.full_lock_ref = child.full_lock_ref @@ -909,7 +896,7 @@ class SWARadixCache(BasePrefixCache): child.key = child.key[split_len:] assert len(child.key) > 0, f"child.key should not be empty" child.value = child.value[split_len:].clone() - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node # insert the new node and child into the lru lists, insert # parent first so that parent is after child in the lru list @@ -938,7 +925,7 @@ class SWARadixCache(BasePrefixCache): if len(key) == 0: return 0 - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children.keys(): @@ -947,7 +934,7 @@ class SWARadixCache(BasePrefixCache): self.full_lru_list.reset_node_mru(node) if not node.swa_tombstone: self.swa_lru_list.reset_node_mru(node) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) if prefix_len < len(node.key): new_node = self._split_node(node.key, node, prefix_len) @@ -1002,7 +989,7 @@ class SWARadixCache(BasePrefixCache): value = value[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) if len(key): # Layout: |--- total_prefix_length ---|--- len(key) ---| @@ -1055,7 +1042,7 @@ class SWARadixCache(BasePrefixCache): new_node.key = key new_node.value = value.clone() new_node.swa_tombstone = swa_tombstone - parent.children[self.get_child_key_fn(key)] = new_node + parent.children[key.child_key(self.page_size)] = new_node self.full_lru_list.insert_mru(new_node) self.full_evictable_size_ += len(value) if not swa_tombstone: @@ -1091,7 +1078,7 @@ class SWARadixCache(BasePrefixCache): not node.swa_tombstone ), f"Invariant violated: leaf node is a tombstone, {node.id=}" assert len(node.children) == 0, f"leaf node has children, {node.id=}" - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" self.full_evictable_size_ -= len(node.key) @@ -1107,7 +1094,7 @@ class SWARadixCache(BasePrefixCache): node.swa_tombstone ), f"Deleting a unexpected non-tombstone leaf node, {node.id=}" assert len(node.children) == 0, f"leaf node has children, {node.id=}" - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node, f"parent does not have child key, {key}" @@ -1152,9 +1139,9 @@ class SWARadixCache(BasePrefixCache): for key, child in current_node.children.items(): stack.append((child, current_indent + 2)) - assert key == self.get_child_key_fn( - child.key - ), f"{key=}, {self.get_child_key_fn(child.key)=}" + assert key == child.key.child_key( + self.page_size + ), f"{key=}, {child.key.child_key(self.page_size)=}" def _total_size_helper(self) -> Tuple[int, int]: total_size = 0 diff --git a/python/sglang/srt/mem_cache/unified_cache_components/README.md b/python/sglang/srt/mem_cache/unified_cache_components/README.md index c698b542a..cf80b76d4 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/README.md +++ b/python/sglang/srt/mem_cache/unified_cache_components/README.md @@ -93,7 +93,7 @@ Find the longest cached prefix for a token sequence. **Algorithm detail:** 1. Calls `create_match_validator()` once per component — returns a stateful closure (e.g., SWA tracks accumulated window length) -2. Walks tree edges via `key_match_fn`; at each node, calls all validator closures — the match boundary is only advanced when **all** validators return `True` +2. Walks tree edges via `RadixKey.match()`; at each node, calls all validator closures — the match boundary is only advanced when **all** validators return `True` 3. If match ends mid-node, calls `_split_node` → triggers `redistribute_on_node_split()` per component 4. Post-match (`_match_post_processor`): - Promotes matched path to MRU in each component's LRU via `node_has_component_data()` as filter diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 0a1a83527..bdaf77fc6 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -20,12 +20,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchPrefixParams, MatchResult, ) -from sglang.srt.mem_cache.radix_cache import ( - RadixKey, - _key_match_page_size1, - _key_match_paged, - get_child_key, -) +from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.unified_cache_components import ( _NUM_COMPONENT_TYPES, BASE_COMPONENT_TYPE, @@ -188,13 +183,6 @@ class UnifiedRadixCache(BasePrefixCache): if params.enable_metrics: self.init_metrics_collector() - if self.page_size == 1: - self.key_match_fn = _key_match_page_size1 - self.get_child_key_fn = get_child_key - else: - self.key_match_fn = partial(_key_match_paged, page_size=self.page_size) - self.get_child_key_fn = partial(get_child_key, page_size=self.page_size) - assert params.tree_components is not None self.tree_components = tuple(params.tree_components) self.components: dict[ComponentType, TreeComponent] = { @@ -484,7 +472,7 @@ class UnifiedRadixCache(BasePrefixCache): Not used yet; reserved for future read-only match operations.""" node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value: list[torch.Tensor] = [] best_value_len = 0 best_node = node @@ -500,7 +488,7 @@ class UnifiedRadixCache(BasePrefixCache): while len(key) > 0 and child_key in node.children: child = node.children[child_key] - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): # Read-only: do not split, ignore partial match and stop break @@ -509,14 +497,14 @@ class UnifiedRadixCache(BasePrefixCache): _update_best_if_valid(node) key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) return value, best_node, best_value_len def _match_prefix_helper( self, key: RadixKey ) -> tuple[list[torch.Tensor], UnifiedTreeNode, int]: node = self.root_node - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) value: list[torch.Tensor] = [] best_value_len = 0 best_node = node @@ -532,7 +520,7 @@ class UnifiedRadixCache(BasePrefixCache): while len(key) > 0 and child_key in node.children: child = node.children[child_key] - prefix_len = self.key_match_fn(child.key, key) + prefix_len = child.key.match(key, page_size=self.page_size) if prefix_len < len(child.key): node = self._split_node(child.key, child, prefix_len) value.append(node.component_data[BASE_COMPONENT_TYPE].value) @@ -543,7 +531,7 @@ class UnifiedRadixCache(BasePrefixCache): _update_best_if_valid(node) key = key[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) return value, best_node, best_value_len def _match_post_processor( @@ -587,7 +575,7 @@ class UnifiedRadixCache(BasePrefixCache): self, key: RadixKey, child: UnifiedTreeNode, split_len: int ) -> UnifiedTreeNode: new_node = UnifiedTreeNode(self.tree_components) - new_node.children = {self.get_child_key_fn(key[split_len:]): child} + new_node.children = {key[split_len:].child_key(self.page_size): child} new_node.parent = child.parent new_node.key = child.key[:split_len] new_node.component_data[BASE_COMPONENT_TYPE].value = ( @@ -604,7 +592,7 @@ class UnifiedRadixCache(BasePrefixCache): for component in self._components_tuple: component.redistribute_on_node_split(new_parent=new_node, child=child) - new_node.parent.children[self.get_child_key_fn(key)] = new_node + new_node.parent.children[key.child_key(self.page_size)] = new_node self._for_each_component_lru(new_node, UnifiedLRUList.insert_mru) self._for_each_component_lru(child, UnifiedLRUList.insert_mru) @@ -626,7 +614,7 @@ class UnifiedRadixCache(BasePrefixCache): new_node.parent = parent new_node.key = key new_node.component_data[BASE_COMPONENT_TYPE].value = value.clone() - parent.children[self.get_child_key_fn(key)] = new_node + parent.children[key.child_key(self.page_size)] = new_node self.lru_lists[BASE_COMPONENT_TYPE].insert_mru(new_node) self.component_evictable_size_[BASE_COMPONENT_TYPE] += len(value) return new_node @@ -642,12 +630,12 @@ class UnifiedRadixCache(BasePrefixCache): if len(key) == 0: return InsertResult(prefix_len=0, mamba_exist=True) - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) total_prefix_length = 0 while len(key) > 0 and child_key in node.children: node = node.children[child_key] self._touch_node(node) - prefix_len = self.key_match_fn(node.key, key) + prefix_len = node.key.match(key, page_size=self.page_size) if prefix_len < len(node.key): node = self._split_node(node.key, node, prefix_len) @@ -674,7 +662,7 @@ class UnifiedRadixCache(BasePrefixCache): key = key[prefix_len:] value = value[prefix_len:] if len(key): - child_key = self.get_child_key_fn(key) + child_key = key.child_key(self.page_size) is_new_leaf = False # Create new leaf for remaining suffix @@ -737,7 +725,7 @@ class UnifiedRadixCache(BasePrefixCache): self._iteratively_delete_tombstone_leaf(node, tracker) def _remove_leaf_from_parent(self, node: UnifiedTreeNode): - key = self.get_child_key_fn(node.key) + key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index eab8f6bfc..72d9da236 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -632,7 +632,7 @@ class UnifiedRadixCacheSuite: self.skipTest("page_size > 1 only") tree, _, _ = build_fixture(self.cfg) key = RadixKey(self._make_seq(1, 1)) - child_key = tree.get_child_key_fn(key) + child_key = key.child_key(tree.page_size) self.assertIsInstance(child_key, tuple) def test_paged_match_truncates_unaligned_key(self):