Make RadixKey slicing zero-copy to fix quadratic match_prefix
RadixKey.__getitem__ copied the token list on every slice, and the match/insert tree walks re-slice the remaining key at every node hop, making a single match O(len * hops) — quadratic for long cached prefixes. Slices now return O(1) offset-based views over a shared backing list; the key-match and child-key functions index the backing list directly so views are never materialized on the hot path. Keys stored in tree nodes are compacted at every store site (same cost as the old copying slices), so lock-ref walks, eviction, splits, and controller-thread reads never observe a key pinning a transient backing list. Node-key compactness is enforced by a tree-walk test. Microbenchmark (64K-token full hit, page_size=64): 16-node path: 1.86 -> 0.87 ms/match (2.1x) 128-node path: 9.80 -> 1.19 ms/match (8.2x), insert re-walk 6.9x Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -730,7 +730,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
child_key = self.get_child_key_fn(key)
|
||||
new_node = TreeNode()
|
||||
new_node.parent = parent
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = value.clone()
|
||||
new_node.mamba_value = mamba_value
|
||||
self.full_lru_list.insert_mru(new_node)
|
||||
@@ -933,7 +933,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
new_node.mamba_value = None
|
||||
new_node.full_lock_ref = child.full_lock_ref
|
||||
new_node.mamba_lock_ref = 0
|
||||
new_node.key = child.key[:split_len]
|
||||
new_node.key = child.key[:split_len].compacted()
|
||||
|
||||
if child.backuped:
|
||||
new_node.host_value = child.host_value[:split_len].clone()
|
||||
@@ -947,7 +947,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
if child.mamba_value is not None:
|
||||
self.mamba_lru_list.remove_node(child)
|
||||
child.parent = new_node
|
||||
child.key = child.key[split_len:]
|
||||
child.key = child.key[split_len:].compacted()
|
||||
new_node.parent.children[self.get_child_key_fn(key)] = new_node
|
||||
if child.mamba_value is not None:
|
||||
self.mamba_lru_list.insert_mru(child)
|
||||
@@ -1649,7 +1649,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
if len(key):
|
||||
new_node = TreeNode()
|
||||
new_node.parent = node
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = None
|
||||
new_node.mamba_value = None
|
||||
new_node.host_value = host_value.clone()
|
||||
|
||||
@@ -4087,7 +4087,7 @@ class HiRadixCache(RadixCache):
|
||||
if len(key):
|
||||
new_node = TreeNode(priority=node.priority)
|
||||
new_node.parent = node
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = None
|
||||
new_node.host_value = host_value.clone()
|
||||
new_node.hash_value = hash_value
|
||||
@@ -4179,7 +4179,7 @@ class HiRadixCache(RadixCache):
|
||||
# If child is pinned, new parent inherits a host_ref_counter hold
|
||||
if child.pin_expiry > 0:
|
||||
new_node.host_ref_counter += 1
|
||||
new_node.key = child.key[:split_len]
|
||||
new_node.key = child.key[:split_len].compacted()
|
||||
new_node.hit_count = child.hit_count
|
||||
|
||||
# split value and host value if exists
|
||||
@@ -4203,7 +4203,7 @@ class HiRadixCache(RadixCache):
|
||||
child.hash_value, split_len, self.page_size
|
||||
)
|
||||
child.parent = new_node
|
||||
child.key = child.key[split_len:]
|
||||
child.key = child.key[split_len:].compacted()
|
||||
new_node.parent.children[self.get_child_key_fn(key)] = new_node
|
||||
|
||||
return new_node
|
||||
@@ -4330,7 +4330,7 @@ class HiRadixCache(RadixCache):
|
||||
priority=priority,
|
||||
)
|
||||
new_node.parent = node
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = value.clone()
|
||||
node.children[child_key] = new_node
|
||||
self.evictable_size_ += self._node_device_resident_len(new_node)
|
||||
|
||||
@@ -1062,7 +1062,7 @@ class MambaRadixCache(BasePrefixCache):
|
||||
new_node.mamba_value = None # mamba cache can not be split
|
||||
new_node.full_lock_ref = child.full_lock_ref
|
||||
new_node.mamba_lock_ref = 0
|
||||
new_node.key = child.key[:split_len]
|
||||
new_node.key = child.key[:split_len].compacted()
|
||||
new_node.value = child.value[:split_len].clone()
|
||||
|
||||
# child time should be later than parent's time for mamba tombstone
|
||||
@@ -1072,7 +1072,7 @@ class MambaRadixCache(BasePrefixCache):
|
||||
if child.mamba_value is not None:
|
||||
self.mamba_lru_list.remove_node(child)
|
||||
child.parent = new_node
|
||||
child.key = child.key[split_len:]
|
||||
child.key = child.key[split_len:].compacted()
|
||||
child.value = child.value[split_len:].clone()
|
||||
new_node.parent.children[self.get_child_key_fn(key)] = new_node
|
||||
|
||||
@@ -1134,7 +1134,7 @@ class MambaRadixCache(BasePrefixCache):
|
||||
if len(key):
|
||||
new_node = TreeNode()
|
||||
new_node.parent = node
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = value.clone()
|
||||
new_node.mamba_value = mamba_value
|
||||
self.full_lru_list.insert_mru(new_node)
|
||||
|
||||
@@ -68,33 +68,87 @@ if TYPE_CHECKING:
|
||||
|
||||
|
||||
class RadixKey:
|
||||
"""A sequence of token ids used as a radix tree key.
|
||||
|
||||
Slicing returns an O(1) view sharing the underlying token list instead of
|
||||
copying it. The tree walks in ``match_prefix``/``insert`` re-slice the
|
||||
remaining key at every node hop, so copying slices makes a single match
|
||||
O(len * hops) — quadratic for long cached prefixes. Views keep it linear.
|
||||
|
||||
Views keep the underlying list alive, so keys stored in long-lived tree
|
||||
nodes must be compacted first (see ``compacted()``); only transient keys
|
||||
inside a tree walk should remain views.
|
||||
"""
|
||||
|
||||
__slots__ = ("_tokens", "_start", "_stop", "extra_key", "is_bigram")
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token_ids: List[int],
|
||||
extra_key: Optional[str] = None,
|
||||
is_bigram: bool = False,
|
||||
):
|
||||
# token ids sequence
|
||||
self.token_ids = token_ids
|
||||
# token ids sequence (full backing list; this key covers
|
||||
# [_start, _stop) of it)
|
||||
self._tokens = token_ids
|
||||
self._start = 0
|
||||
self._stop = len(token_ids)
|
||||
# extra key (e.g. lora_id, cache_salt)
|
||||
self.extra_key = extra_key
|
||||
# is bigram key
|
||||
self.is_bigram = is_bigram
|
||||
|
||||
@classmethod
|
||||
def _view(cls, base: "RadixKey", start: int, stop: int) -> "RadixKey":
|
||||
obj = cls.__new__(cls)
|
||||
obj._tokens = base._tokens
|
||||
obj._start = start
|
||||
obj._stop = stop
|
||||
obj.extra_key = base.extra_key
|
||||
obj.is_bigram = base.is_bigram
|
||||
return obj
|
||||
|
||||
@property
|
||||
def token_ids(self) -> List[int]:
|
||||
if self._start == 0 and self._stop == len(self._tokens):
|
||||
return self._tokens
|
||||
return self._tokens[self._start : self._stop]
|
||||
|
||||
@token_ids.setter
|
||||
def token_ids(self, value: List[int]) -> None:
|
||||
self._tokens = value
|
||||
self._start = 0
|
||||
self._stop = len(value)
|
||||
|
||||
def compacted(self) -> "RadixKey":
|
||||
"""Return an equivalent key that does not pin a larger backing list."""
|
||||
if self._start == 0 and self._stop == len(self._tokens):
|
||||
return self
|
||||
return RadixKey(
|
||||
self._tokens[self._start : self._stop], self.extra_key, self.is_bigram
|
||||
)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.token_ids)
|
||||
return self._stop - self._start
|
||||
|
||||
def __iter__(self) -> Iterator[int]:
|
||||
return iter(self.token_ids)
|
||||
# NOTE: map over an index range instead of islice — islice skips the
|
||||
# first `_start` elements one by one, which is O(_start) per call.
|
||||
return map(self._tokens.__getitem__, range(self._start, self._stop))
|
||||
|
||||
def __getitem__(self, idx: Union[int, slice]) -> "RadixKey":
|
||||
if isinstance(idx, slice):
|
||||
return RadixKey(self.token_ids[idx], self.extra_key)
|
||||
return RadixKey([self.token_ids[idx]], self.extra_key)
|
||||
start, stop, step = idx.indices(len(self))
|
||||
if step != 1:
|
||||
raise ValueError(f"RadixKey does not support strided slicing: {idx}")
|
||||
return RadixKey._view(self, self._start + start, self._start + stop)
|
||||
if idx < 0:
|
||||
idx += len(self)
|
||||
return RadixKey([self._tokens[self._start + idx]], self.extra_key)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
preview = self.token_ids[:10]
|
||||
return f"RadixKey(extra_key={self.extra_key!r}, token_ids={preview}{'...' if len(self.token_ids) > 10 else ''})"
|
||||
return f"RadixKey(extra_key={self.extra_key!r}, token_ids={preview}{'...' if len(self) > 10 else ''})"
|
||||
|
||||
|
||||
def maybe_bigram_convert(
|
||||
@@ -207,22 +261,27 @@ def _check_extra_key(key0: RadixKey, key1: RadixKey):
|
||||
|
||||
def _key_match_page_size1(key0: RadixKey, key1: RadixKey):
|
||||
_check_extra_key(key0, key1)
|
||||
# Index the backing lists directly so that view keys (nonzero _start) are
|
||||
# matched without materializing a copy of their suffix.
|
||||
t0, s0 = key0._tokens, key0._start
|
||||
t1, s1 = key1._tokens, key1._start
|
||||
n = min(len(key0), len(key1))
|
||||
i = 0
|
||||
for k0, k1 in zip(key0.token_ids, key1.token_ids):
|
||||
if k0 != k1:
|
||||
break
|
||||
while i < n and t0[s0 + i] == t1[s1 + i]:
|
||||
i += 1
|
||||
return i
|
||||
|
||||
|
||||
def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
|
||||
_check_extra_key(key0, key1)
|
||||
t0, s0 = key0._tokens, key0._start
|
||||
t1, s1 = key1._tokens, key1._start
|
||||
min_len = min(len(key0), len(key1))
|
||||
|
||||
i = 0
|
||||
while i < min_len:
|
||||
step = min(page_size, min_len - i)
|
||||
if key0.token_ids[i : i + step] != key1.token_ids[i : i + step]:
|
||||
if t0[s0 + i : s0 + i + step] != t1[s1 + i : s1 + i + step]:
|
||||
break
|
||||
i += step
|
||||
|
||||
@@ -231,9 +290,9 @@ def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
|
||||
|
||||
def get_child_key(key: RadixKey, page_size: int = 1):
|
||||
if page_size == 1:
|
||||
plain_key = key.token_ids[0]
|
||||
plain_key = key._tokens[key._start]
|
||||
else:
|
||||
plain_key = tuple(key.token_ids[:page_size])
|
||||
plain_key = tuple(key._tokens[key._start : key._start + page_size])
|
||||
if key.extra_key is None:
|
||||
return plain_key
|
||||
else:
|
||||
@@ -892,10 +951,12 @@ class RadixCache(BasePrefixCache):
|
||||
new_node.children = {self.get_child_key_fn(key[split_len:]): child}
|
||||
new_node.parent = child.parent
|
||||
new_node.lock_ref = child.lock_ref
|
||||
new_node.key = child.key[:split_len]
|
||||
# Node keys must be compacted: a view would pin the backing list of the
|
||||
# (possibly transient) key it was sliced from for the node's lifetime.
|
||||
new_node.key = child.key[:split_len].compacted()
|
||||
new_node.value = child.value[:split_len].clone()
|
||||
child.parent = new_node
|
||||
child.key = child.key[split_len:]
|
||||
child.key = child.key[split_len:].compacted()
|
||||
child.value = child.value[split_len:].clone()
|
||||
new_node.parent.children[self.get_child_key_fn(key)] = new_node
|
||||
|
||||
@@ -957,7 +1018,7 @@ class RadixCache(BasePrefixCache):
|
||||
if len(key):
|
||||
new_node = TreeNode(priority=priority)
|
||||
new_node.parent = node
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = value.clone()
|
||||
self._inc_hit_count(new_node, chunked)
|
||||
node.children[child_key] = new_node
|
||||
|
||||
@@ -191,7 +191,7 @@ class LMCRadixCache(RadixCache):
|
||||
new_node = TreeNode(priority=last_node.priority)
|
||||
start = value.numel()
|
||||
end = start + fetched
|
||||
new_node.key = key[start:end]
|
||||
new_node.key = key[start:end].compacted()
|
||||
new_node.value = token_slots[:fetched]
|
||||
new_node.parent = last_node
|
||||
last_node.children[self.get_child_key_fn(new_node.key)] = new_node
|
||||
|
||||
@@ -899,7 +899,7 @@ class SWARadixCache(BasePrefixCache):
|
||||
new_node.swa_tombstone = child.swa_tombstone
|
||||
new_node.full_lock_ref = child.full_lock_ref
|
||||
new_node.swa_lock_ref = child.swa_lock_ref
|
||||
new_node.key = child.key[:split_len]
|
||||
new_node.key = child.key[:split_len].compacted()
|
||||
assert len(new_node.key) > 0, f"new_node.key should not be empty"
|
||||
new_node.value = child.value[:split_len].clone()
|
||||
# parent inherits the swa_uuid from child for swa lock ref
|
||||
@@ -913,7 +913,7 @@ class SWARadixCache(BasePrefixCache):
|
||||
if not new_node.swa_tombstone:
|
||||
self.swa_lru_list.remove_node(child)
|
||||
child.parent = new_node
|
||||
child.key = child.key[split_len:]
|
||||
child.key = child.key[split_len:].compacted()
|
||||
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
|
||||
@@ -1039,7 +1039,7 @@ class SWARadixCache(BasePrefixCache):
|
||||
assert len(key) > 0, f"key should not be empty"
|
||||
new_node = TreeNode()
|
||||
new_node.parent = parent
|
||||
new_node.key = key
|
||||
new_node.key = key.compacted()
|
||||
new_node.value = value.clone()
|
||||
new_node.swa_tombstone = swa_tombstone
|
||||
parent.children[self.get_child_key_fn(key)] = new_node
|
||||
|
||||
Reference in New Issue
Block a user