diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index e6d5ba594..dd5342f4b 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -180,10 +180,18 @@ class RadixKey: def match(self, other: RadixKey, page_size: int = 1) -> int: """Logical-unit prefix length shared with ``other``. Result is rounded down to ``page_size``.""" + return self.match_at(other, offset=0, page_size=page_size) + + def match_at(self, other: RadixKey, offset: int, page_size: int = 1) -> int: + """Match without slicing while preserving bigram boundaries and limit semantics.""" self._check_compatible(other) + if self.is_bigram != other.is_bigram: + raise ValueError("RadixKey operations require matching bigram modes") + if offset < 0 or offset > len(other): + raise IndexError(f"RadixKey offset out of range: {offset}") t0, t1 = self.token_ids, other.token_ids assert type(t0) is type(t1), (type(t0), type(t1)) - n = min(len(t0), len(t1)) + n = min(self._raw_len(), other._raw_len() - offset) # Exponential search for the first diverging token: gallop in doubling # windows (one C-level slice compare each), then binary-search the window @@ -193,10 +201,10 @@ class RadixKey: step = 1 while lo < n: hi = lo + step if lo + step < n else n - if t0[lo:hi] != t1[lo:hi]: + if t0[lo:hi] != t1[offset + lo : offset + hi]: while hi - lo > 1: mid = (lo + hi) // 2 - if t0[lo:mid] == t1[lo:mid]: + if t0[lo:mid] == t1[offset + lo : offset + mid]: lo = mid else: hi = mid @@ -206,24 +214,37 @@ class RadixKey: step *= 2 if self.is_bigram: - matched = max(0, min(matched_tokens - 1, len(self), len(other))) + matched = max(0, min(matched_tokens - 1, len(self), len(other) - offset)) return (matched // page_size) * page_size if page_size > 1 else matched - matched_tokens = min(matched_tokens, len(self), len(other)) + matched_tokens = min(matched_tokens, len(self), len(other) - offset) if page_size == 1: return matched_tokens return (matched_tokens // page_size) * page_size def child_key(self, page_size: int = 1): """Hashable dict-key for the first ``page_size`` logical units, namespaced by ``extra_key``.""" + return self.child_key_at(offset=0, page_size=page_size) + + def child_key_at(self, offset: int, page_size: int = 1): + """Hashable child key at ``offset`` without slicing token storage.""" + if offset < 0 or offset + page_size > len(self): + raise IndexError( + f"RadixKey child range out of bounds: offset={offset}, " + f"page_size={page_size}, len={len(self)}" + ) t = self.token_ids if self.is_bigram: if page_size == 1: - plain = (t[0], t[1]) + plain = (t[offset], t[offset + 1]) else: - plain = tuple((t[j], t[j + 1]) for j in range(page_size)) + plain = tuple( + (t[j], t[j + 1]) for j in range(offset, offset + page_size) + ) else: - plain = t[0] if page_size == 1 else tuple(t[:page_size]) + plain = ( + t[offset] if page_size == 1 else tuple(t[offset : offset + page_size]) + ) if self.cache_salt is not None: return ((self.extra_key, self.cache_salt), plain) return plain if self.extra_key is None else (self.extra_key, plain) 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 e25f7bc95..7dc939be4 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 @@ -737,7 +737,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): # nodes can also match, so we separately track the best device-resident # match for scheduler prefix indices and locking. node = self.root_node - child_key = key.child_key(self.page_size) + key_offset = 0 + child_key = key.child_key_at(key_offset, self.page_size) value: list[torch.Tensor] = [] best_match_node = node best_match_device_node = node @@ -778,14 +779,14 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): best_match_device_value_len = len(value) best_match_device_node = node - while len(key) > 0 and child_key in node.children: + while key_offset < len(key) and child_key in node.children: child = node.children[child_key] # HiCache: dead node (evicted + not backuped) — stop traversal if child.evicted and not child.backuped: break - prefix_len = child.key.match(key, page_size=self.page_size) + prefix_len = child.key.match_at(key, key_offset, page_size=self.page_size) full_kv_hit_length += prefix_len if prefix_len < len(child.key): node, action = self._split_node(child.key, child, prefix_len) @@ -798,9 +799,9 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): value.append(child.component_data[BASE_COMPONENT_TYPE].value) node = child _update_best_if_valid(node) - key = key[prefix_len:] - if len(key): - child_key = key.child_key(self.page_size) + key_offset += prefix_len + if key_offset < len(key): + child_key = key.child_key_at(key_offset, self.page_size) return ( value,