Refactors radix cache for extra key support (#10317)

Signed-off-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
Xinyuan Tong
2025-09-21 11:16:16 -07:00
committed by GitHub
parent fc3e542009
commit 12d6cf18f0
13 changed files with 821 additions and 574 deletions

View File

@@ -23,7 +23,7 @@ import heapq
import time
from collections import defaultdict
from functools import partial
from typing import TYPE_CHECKING, List, Optional
from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Union
import torch
@@ -41,6 +41,30 @@ if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
class RadixKey:
def __init__(self, token_ids: List[int], extra_key: Optional[str] = None):
# token ids sequence
self.token_ids = token_ids
# extra key (e.g. lora_id, cache_salt)
self.extra_key = extra_key
def __len__(self) -> int:
return len(self.token_ids)
def __iter__(self) -> Iterator[int]:
return iter(self.token_ids)
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)
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 ''})"
class TreeNode:
counter = 0
@@ -48,7 +72,7 @@ class TreeNode:
def __init__(self, id: Optional[int] = None):
self.children = defaultdict(TreeNode)
self.parent: TreeNode = None
self.key: List[int] = None
self.key: RadixKey = None
self.value: Optional[torch.Tensor] = None
self.lock_ref = 0
self.last_access_time = time.monotonic()
@@ -94,27 +118,47 @@ class TreeNode:
return self.last_access_time < other.last_access_time
def _key_match_page_size1(key0: List, key1: List):
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)
i = 0
for k0, k1 in zip(key0, key1):
for k0, k1 in zip(key0.token_ids, key1.token_ids):
if k0 != k1:
break
i += 1
return i
def _key_match_paged(key0: List, key1: List, page_size: int):
def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
_check_extra_key(key0, key1)
min_len = min(len(key0), len(key1))
i = 0
while i < min_len:
if key0[i : i + page_size] != key1[i : i + page_size]:
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 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)
class RadixCache(BasePrefixCache):
def __init__(
self,
@@ -139,10 +183,10 @@ class RadixCache(BasePrefixCache):
if self.page_size == 1:
self.key_match_fn = _key_match_page_size1
self.get_child_key_fn = lambda key: key[0]
self.get_child_key_fn = get_child_key
else:
self.key_match_fn = partial(_key_match_paged, page_size=page_size)
self.get_child_key_fn = lambda key: tuple(key[:page_size])
self.get_child_key_fn = partial(get_child_key, page_size=page_size)
if eviction_policy.lower() == "lru":
self.eviction_strategy: EvictionStrategy = LRUStrategy()
@@ -158,7 +202,7 @@ class RadixCache(BasePrefixCache):
def reset(self):
self.root_node = TreeNode()
self.root_node.key = []
self.root_node.key = RadixKey(token_ids=[], extra_key=None)
self.root_node.value = []
self.root_node.host_value = []
self.root_node.lock_ref = 1
@@ -166,16 +210,43 @@ class RadixCache(BasePrefixCache):
self.protected_size_ = 0
self._record_all_cleared_event()
def match_prefix(self, key: List[int], **kwargs) -> MatchResult:
"""Find the matching prefix from the radix tree.
def match_prefix(self, key: RadixKey, **kwargs) -> MatchResult:
"""Find the longest cached prefix of ``key`` in the radix tree.
The logical namespace for prefix matching is determined by both the
token id sequence and the optional ``extra_key`` carried by ``RadixKey``.
Entries that share identical leading token ids but have *different*
``extra_key`` values are intentionally kept disjoint and never share
prefix nodes. This is useful to:
* Isolate KV cache lines for different LoRA / adapter IDs.
* Separate requests that intentionally should not share state (e.g.,
different sampling salt, cache version, or retrieval augmentation
context) by supplying a distinct ``extra_key``.
Args:
key: A list of token IDs to find a matching prefix.
key (RadixKey): The lookup key containing a list of token ids and an
optional ``extra_key`` namespace tag. If ``page_size > 1`` the
length is internally truncated to a multiple of ``page_size``
before matching. Passing an empty key returns an empty result
with the root as the last node.
**kwargs: Reserved for future extensions (ignored currently).
Returns:
A tuple of a tensor of matching prefix token IDs and
the last node that contains the prefix values. Note that
this API can modify the internal state of the Radix tree.
The last node create a new child if the prefix is shorter
than the last node's value.
MatchResult: ``device_indices`` is a 1-D ``torch.int64`` tensor of
the concatenated KV cache indices corresponding to the longest
cached prefix (may be length 0). ``last_device_node`` and
``last_host_node`` (currently the same) are the tree node objects
representing the terminal node of the matched prefix. This method
may mutate internal structure by splitting an existing node if the
match ends inside a stored segment.
Internal updates:
* Refreshes access metadata (timestamps) used by the
configured eviction strategy.
* If the lookup ends inside a stored segment the node is split once
to expose a precise boundary; this structural refinement improves
subsequent match efficiency and does not duplicate data.
"""
if self.disable or len(key) == 0:
return MatchResult(
@@ -203,12 +274,12 @@ class RadixCache(BasePrefixCache):
last_host_node=last_node,
)
def insert(self, key: List, value=None, chunked=False):
def insert(self, key: RadixKey, value=None, chunked=False):
if self.disable:
return 0
if value is None:
value = [x for x in key]
value = torch.tensor(key.token_ids, dtype=torch.int64)
return self._insert_helper(self.root_node, key, value)
def cache_finished_req(self, req: Req):
@@ -238,7 +309,8 @@ class RadixCache(BasePrefixCache):
# Radix Cache takes one ref in memory pool
new_prefix_len = self.insert(
token_ids[:page_aligned_len], page_aligned_kv_indices
RadixKey(token_ids[:page_aligned_len], req.extra_key),
page_aligned_kv_indices,
)
self.token_to_kv_pool_allocator.free(
kv_indices[len(req.prefix_indices) : new_prefix_len]
@@ -270,14 +342,18 @@ class RadixCache(BasePrefixCache):
# Radix Cache takes one ref in memory pool
new_prefix_len = self.insert(
page_aligned_token_ids, page_aligned_kv_indices, chunked=chunked
RadixKey(page_aligned_token_ids, req.extra_key),
page_aligned_kv_indices,
chunked=chunked,
)
self.token_to_kv_pool_allocator.free(
kv_indices[len(req.prefix_indices) : new_prefix_len]
)
# The prefix indices could be updated, reuse it
new_indices, new_last_node, _, _ = self.match_prefix(page_aligned_token_ids)
new_indices, new_last_node, _, _ = self.match_prefix(
RadixKey(token_ids=page_aligned_token_ids, extra_key=req.extra_key)
)
self.req_to_token_pool.write(
(req.req_pool_idx, slice(len(req.prefix_indices), len(new_indices))),
new_indices[len(req.prefix_indices) :],
@@ -379,7 +455,7 @@ class RadixCache(BasePrefixCache):
##### Internal Helper Functions #####
def _match_prefix_helper(self, node: TreeNode, key: List):
def _match_prefix_helper(self, node: TreeNode, key: RadixKey):
node.last_access_time = time.monotonic()
child_key = self.get_child_key_fn(key)
@@ -404,7 +480,7 @@ class RadixCache(BasePrefixCache):
return value, node
def _split_node(self, key, child: TreeNode, split_len: int):
def _split_node(self, key: RadixKey, child: TreeNode, split_len: int):
# new_node -> child
self._record_remove_event(child)
new_node = TreeNode()
@@ -423,7 +499,7 @@ class RadixCache(BasePrefixCache):
return new_node
def _insert_helper(self, node: TreeNode, key: List, value):
def _insert_helper(self, node: TreeNode, key: RadixKey, value):
node.last_access_time = time.monotonic()
if len(key) == 0:
return 0
@@ -464,7 +540,7 @@ class RadixCache(BasePrefixCache):
print(
" " * current_indent,
len(current_node.key),
current_node.key[:10],
current_node.key.token_ids[:10],
f"r={current_node.lock_ref}",
)
for key, child in current_node.children.items():
@@ -516,11 +592,11 @@ class RadixCache(BasePrefixCache):
last_page_start = (
(len(node.parent.key) - 1) // self.page_size
) * self.page_size
parent_parent_tokens = node.parent.key[last_page_start:]
parent_parent_tokens = node.parent.key.token_ids[last_page_start:]
parent_block_hash = hash(tuple(parent_parent_tokens))
for start in range(0, len(node.key), self.page_size):
page_tokens = node.key[start : start + self.page_size]
page_tokens = node.key.token_ids[start : start + self.page_size]
if not page_tokens:
continue
@@ -543,7 +619,7 @@ class RadixCache(BasePrefixCache):
# One BlockRemoved per chunk.
if self.enable_kv_cache_events:
for start in range(0, len(node.key), self.page_size):
page_tokens = node.key[start : start + self.page_size]
page_tokens = node.key.token_ids[start : start + self.page_size]
if not page_tokens:
continue
block_hash = hash(tuple(page_tokens))
@@ -569,19 +645,12 @@ class RadixCache(BasePrefixCache):
if __name__ == "__main__":
tree = RadixCache(None, None, page_size=1, disable=False)
tree.insert("Hello")
tree.insert("Hello")
tree.insert("Hello_L.A.!")
# tree.insert("Hello_world! Happy")
# tree.insert("I love you!")
# Example token id sequences (as lists of ints)
tree.insert(RadixKey(token_ids=[1, 2, 3], extra_key=None))
tree.insert(RadixKey(token_ids=[1, 2, 3], extra_key=None))
tree.insert(RadixKey(token_ids=[1, 2, 4, 5], extra_key=None))
tree.insert(RadixKey(token_ids=[1, 2, 4, 5, 6, 7], extra_key=None))
tree.insert(RadixKey(token_ids=[8, 9, 10, 11, 12], extra_key=None))
tree.pretty_print()
# print(tree.match_prefix("I love you! aha"))
# def evict_callback(x):
# print("evict", x)
# return len(x)
# tree.evict(5, evict_callback)
# tree.evict(10, evict_callback)
# tree.pretty_print()
print(tree.match_prefix(RadixKey(token_ids=[1, 2, 3, 13, 14], extra_key=None)))