[Perf] Make EAGLE bigram key an O(1) view on RadixKey (#23106)

This commit is contained in:
Liangsheng Yin
2026-04-20 12:01:11 -07:00
committed by GitHub
parent 3dc1491c95
commit 8cb957ccff
5 changed files with 185 additions and 104 deletions
+2 -2
View File
@@ -1207,7 +1207,7 @@ class HiRadixCache(RadixCache):
def match_prefix(self, params: MatchPrefixParams):
key = params.key
empty_value = torch.empty((0,), dtype=torch.int64, device=self.device)
key, _ = self.maybe_bigram_convert(key)
key, _ = key.maybe_to_bigram_view(self.is_eagle)
if self.disable or len(key) == 0:
return MatchResult(
device_indices=empty_value,
@@ -1394,7 +1394,7 @@ class HiRadixCache(RadixCache):
if priority is None:
priority = 0
key, value = self.maybe_bigram_convert(key, value)
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
if len(key) == 0:
return InsertResult(prefix_len=0)
+158 -70
View File
@@ -1,7 +1,6 @@
from __future__ import annotations
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.utils import convert_to_bigram_key
"""
Copyright 2023-2024 SGLang Team
@@ -62,58 +61,95 @@ from sglang.srt.mem_cache.evict_policy import (
PriorityStrategy,
SLRUStrategy,
)
from sglang.srt.mem_cache.utils import get_hash_str, hash_str_to_int64
from sglang.srt.mem_cache.utils import hash_str_to_int64
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
class RadixKey:
"""is_bigram=True: token_ids holds raw tokens (N+1 for N bigrams); slices share one boundary token."""
__slots__ = ("token_ids", "extra_key", "is_bigram")
def __init__(
self,
token_ids: List[int],
extra_key: Optional[str] = None,
is_bigram: bool = False,
):
# token ids sequence
# token ids sequence (raw ints in both modes)
self.token_ids = token_ids
# extra key (e.g. lora_id, cache_salt)
self.extra_key = extra_key
# is bigram key
# bigram view over token_ids: length = max(0, len(token_ids) - 1)
self.is_bigram = is_bigram
def __len__(self) -> int:
if self.is_bigram:
n = len(self.token_ids)
return n - 1 if n > 0 else 0
return len(self.token_ids)
def __iter__(self) -> Iterator[int]:
return iter(self.token_ids)
def __iter__(self) -> Iterator:
if self.is_bigram:
t = self.token_ids
for i in range(len(t) - 1):
yield (t[i], t[i + 1])
else:
yield from 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)
# Normalize int -> 1-element slice so the rest handles one shape.
if isinstance(idx, int):
if idx < 0:
idx += len(self)
if idx < 0 or idx >= len(self):
raise IndexError(f"RadixKey index out of range: {idx}")
idx = slice(idx, idx + 1)
start, stop, step = idx.indices(len(self))
if step != 1:
raise ValueError("RadixKey slice step must be 1")
if self.is_bigram:
# bigrams [start, stop) span raw tokens [start, stop + 1);
# empty slice -> empty raw tokens (not a dangling boundary token).
raw = self.token_ids[start : stop + 1] if stop > start else []
return RadixKey(raw, self.extra_key, is_bigram=True)
return RadixKey(self.token_ids[start:stop], 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.token_ids) > 10 else ''}, is_bigram={self.is_bigram})"
def maybe_to_bigram_view(
self,
is_eagle: bool,
value: Optional[torch.Tensor] = None,
) -> Tuple["RadixKey", Optional[torch.Tensor]]:
# O(1): flip the bigram flag instead of materializing a tuple list.
# value is paired with raw tokens and gets truncated to the bigram count.
if is_eagle and not self.is_bigram:
self.is_bigram = True
if value is not None:
value = value[: len(self)]
return self, value
def maybe_bigram_convert(
is_eagle: bool,
key: RadixKey,
value: Optional[torch.Tensor] = None,
) -> Tuple[RadixKey, Optional[torch.Tensor]]:
if is_eagle and not key.is_bigram:
key.token_ids = convert_to_bigram_key(key.token_ids)
key.is_bigram = True
if value is not None:
value = value[: len(key)]
return key, value
def page_align_keys(key: list, page_size: int, is_bigram: bool = False) -> list:
"""Truncate a raw token list so the resulting RadixKey length is page-aligned.
def page_align_keys(key: list, page_size) -> list:
In bigram mode, logical length = len(key) - 1, and we must keep one extra
boundary token so that bigram_count == aligned.
"""
if page_size == 1:
return key
if is_bigram:
logical_len = len(key) - 1 if len(key) > 0 else 0
aligned = logical_len // page_size * page_size
if aligned == 0:
return []
return key[: aligned + 1]
page_aligned_len = len(key) // page_size * page_size
return key[:page_aligned_len]
@@ -190,32 +226,60 @@ def _check_extra_key(key0: RadixKey, key1: RadixKey):
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 k0, k1 in zip(key0.token_ids, key1.token_ids):
if k0 != k1:
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)
min_len = min(len(key0), len(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 page_size == 1:
plain_key = key.token_ids[0]
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:
plain_key = tuple(key.token_ids[:page_size])
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:
@@ -225,36 +289,56 @@ def get_child_key(key: RadixKey, page_size: int = 1):
def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]:
"""Compute SHA256-based hash values for position-aware identification.
Args:
node: The TreeNode to compute hash values for
page_size: The page size for chunking tokens
Returns:
List of SHA256 hex strings, one per page
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.
"""
hash_values = []
# Get parent's last hash value if parent exists
parent_hash = None
if node.parent is not None and node.parent.hash_value is not None:
# Check if parent is root by checking if it has empty key
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
parent_hash = node.parent.hash_value[-1]
# Iterate through node's pages
for start in range(0, len(node.key), page_size):
page_tokens = node.key.token_ids[start : start + page_size]
if not page_tokens:
continue
raw = node.key.token_ids
is_bigram = node.key.is_bigram
logical_len = len(node.key)
# Use SHA256-based chaining via get_hash_str
hash_val = get_hash_str(page_tokens, prior_hash=parent_hash)
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_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]]]:
@@ -366,11 +450,6 @@ class RadixCache(BasePrefixCache):
self.evictable_leaves.clear()
self._record_all_cleared_event()
def maybe_bigram_convert(
self, key: RadixKey, value: Optional[torch.Tensor] = None
) -> Tuple[RadixKey, Optional[torch.Tensor]]:
return maybe_bigram_convert(self.is_eagle, key, value)
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
"""Find the longest cached prefix of ``key`` in the radix tree.
@@ -409,7 +488,7 @@ class RadixCache(BasePrefixCache):
subsequent match efficiency and does not duplicate data.
"""
key = params.key
key, _ = self.maybe_bigram_convert(key)
key, _ = key.maybe_to_bigram_view(self.is_eagle)
def empty_match_result():
return MatchResult(
@@ -453,9 +532,11 @@ class RadixCache(BasePrefixCache):
chunked = params.chunked
if value is None:
value = torch.tensor(key.token_ids, dtype=torch.int64)
# Debug/test fallback: use token ids themselves as values. Truncate
# to the logical key length so bigram mode gets len(key) entries.
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
key, value = self.maybe_bigram_convert(key, value)
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
prefix_len = self._insert_helper(self.root_node, key, value, priority, chunked)
return InsertResult(prefix_len=prefix_len)
@@ -479,11 +560,9 @@ class RadixCache(BasePrefixCache):
req.req_pool_idx, : len(token_ids)
]
# Maybe convert to bigram keys for EAGLE
keys = convert_to_bigram_key(token_ids) if self.is_eagle else token_ids
keys = page_align_keys(keys, self.page_size)
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
values = kv_indices[: len(radix_key)].to(dtype=torch.int64, copy=True)
# Radix Cache takes one ref in memory pool
if is_insert:
@@ -498,11 +577,11 @@ class RadixCache(BasePrefixCache):
)
else:
self.token_to_kv_pool_allocator.free(
kv_indices[req.cache_protected_len : len(keys)]
kv_indices[req.cache_protected_len : len(radix_key)]
)
# free the unaligned tail
self.token_to_kv_pool_allocator.free(kv_indices[len(keys) :])
self.token_to_kv_pool_allocator.free(kv_indices[len(radix_key) :])
# Remove req slot release the cache lock
self.dec_lock_ref(req.last_node)
@@ -517,11 +596,9 @@ class RadixCache(BasePrefixCache):
req.req_pool_idx, : len(token_ids)
]
# Maybe convert to bigram keys for EAGLE
keys = convert_to_bigram_key(token_ids) if self.is_eagle else token_ids
keys = page_align_keys(keys, self.page_size)
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
values = kv_indices[: len(radix_key)].to(dtype=torch.int64, copy=True)
# Radix Cache takes one ref in memory pool
result = self.insert(
@@ -544,7 +621,9 @@ class RadixCache(BasePrefixCache):
match_result.device_indices,
match_result.last_device_node,
)
assert len(new_indices) == len(keys), f"{len(new_indices)=}, {len(keys)=}"
assert len(new_indices) == len(
radix_key
), f"{len(new_indices)=}, {len(radix_key)=}"
self.req_to_token_pool.write(
(req.req_pool_idx, slice(req.cache_protected_len, len(new_indices))),
@@ -846,10 +925,18 @@ class RadixCache(BasePrefixCache):
parent_block_hash = hash_str_to_int64(node.parent.hash_value[-1])
page_index = 0
for start in range(0, len(node.key), self.page_size):
page_tokens = node.key.token_ids[start : start + self.page_size]
if not page_tokens:
logical_len = len(node.key)
is_bigram = node.key.is_bigram
raw = node.key.token_ids
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
# Preserve historical event payload: bigram pages expose tuples.
if is_bigram:
page_tokens = [(raw[j], raw[j + 1]) for j in range(start, end)]
else:
page_tokens = raw[start:end]
block_hash = hash_str_to_int64(node.hash_value[page_index])
@@ -875,9 +962,10 @@ class RadixCache(BasePrefixCache):
node.hash_value = compute_node_hash_values(node, self.page_size)
page_index = 0
for start in range(0, len(node.key), self.page_size):
page_tokens = node.key.token_ids[start : start + self.page_size]
if not page_tokens:
logical_len = len(node.key)
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
block_hash = hash_str_to_int64(node.hash_value[page_index])
+9 -16
View File
@@ -46,7 +46,6 @@ from sglang.srt.mem_cache.radix_cache import (
_key_match_page_size1,
_key_match_paged,
get_child_key,
maybe_bigram_convert,
page_align_keys,
)
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
@@ -432,9 +431,9 @@ class SWARadixCache(BasePrefixCache):
swa_evicted_seqlen = params.swa_evicted_seqlen
if value is None:
value = torch.tensor([x for x in key.token_ids], dtype=torch.int64)
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
key, value = maybe_bigram_convert(self.is_eagle, key, value)
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
prefix_len = self._insert_helper(
self.root_node, key, value, prev_prefix_len, swa_evicted_seqlen
@@ -456,16 +455,11 @@ class SWARadixCache(BasePrefixCache):
req.req_pool_idx, :kv_committed_len
]
# Maybe convert to bigram keys for EAGLE
keys = self.key_convert_fn(token_ids)
keys = page_align_keys(keys, self.page_size)
page_aligned_len = len(keys)
# EAGLE: skip tuple materialization; is_bigram flag gives bigram semantics.
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
page_aligned_len = len(radix_key)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
radix_key = RadixKey(
keys[:page_aligned_len],
req.extra_key,
is_bigram=self.is_eagle,
)
old_prefix_len = req.cache_protected_len
# Radix Cache takes one ref in memory pool
@@ -508,10 +502,9 @@ class SWARadixCache(BasePrefixCache):
req.req_pool_idx, : len(token_ids)
]
keys = self.key_convert_fn(token_ids)
keys = page_align_keys(keys, self.page_size)
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
values = kv_indices[: len(radix_key)].to(dtype=torch.int64, copy=True)
old_prefix_len = req.cache_protected_len
# Radix Cache takes one ref in memory pool
@@ -846,7 +839,7 @@ class SWARadixCache(BasePrefixCache):
def _match_pre_processor(self, params: MatchPrefixParams) -> Optional[RadixKey]:
"""Preprocess the key before matching."""
key = params.key
key, _ = maybe_bigram_convert(self.is_eagle, key)
key, _ = key.maybe_to_bigram_view(self.is_eagle)
if self.disable or len(key) == 0:
return None
@@ -25,7 +25,6 @@ from sglang.srt.mem_cache.radix_cache import (
_key_match_page_size1,
_key_match_paged,
get_child_key,
maybe_bigram_convert,
page_align_keys,
)
from sglang.srt.mem_cache.unified_cache_components import (
@@ -239,7 +238,7 @@ class UnifiedRadixCache(BasePrefixCache):
return result
key = params.key
key, _ = maybe_bigram_convert(self.is_eagle, key)
key, _ = key.maybe_to_bigram_view(self.is_eagle)
if self.disable or len(key) == 0:
return MatchResult(
device_indices=torch.empty(
@@ -264,9 +263,9 @@ class UnifiedRadixCache(BasePrefixCache):
key = params.key
value = params.value
if value is None:
value = torch.tensor([x for x in key.token_ids], dtype=torch.int64)
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
key, value = maybe_bigram_convert(self.is_eagle, key, value)
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
result = self._insert_helper(self.root_node, key, value, params)
return result
@@ -355,12 +354,11 @@ class UnifiedRadixCache(BasePrefixCache):
token_ids = token_ids[:effective_cache_len]
kv_indices = kv_indices[:effective_cache_len]
# Key convert + page align
keys = self.key_convert_fn(token_ids)
keys = page_align_keys(keys, self.page_size)
page_aligned_len = len(keys)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
# Page align on raw tokens; bigram semantics via is_bigram flag.
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
page_aligned_len = len(radix_key)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
insert_params.key = radix_key
insert_params.value = values
@@ -422,12 +420,13 @@ class UnifiedRadixCache(BasePrefixCache):
kv_indices = kv_indices_orig[:effective_cache_len]
# Key convert + page align
keys = self.key_convert_fn(token_ids[:effective_cache_len])
keys = page_align_keys(keys, self.page_size)
page_aligned_len = len(keys)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
# Page align on raw tokens; bigram semantics via is_bigram flag.
keys = page_align_keys(
token_ids[:effective_cache_len], self.page_size, is_bigram=self.is_eagle
)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
page_aligned_len = len(radix_key)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
insert_params.key = radix_key
insert_params.value = values