From 8cb957ccffa23da7ded70f1efdc5d5f5e9b6578a Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Mon, 20 Apr 2026 12:01:11 -0700 Subject: [PATCH] [Perf] Make EAGLE bigram key an O(1) view on `RadixKey` (#23106) --- python/sglang/srt/mem_cache/hiradix_cache.py | 4 +- python/sglang/srt/mem_cache/radix_cache.py | 228 ++++++++++++------ .../sglang/srt/mem_cache/swa_radix_cache.py | 25 +- .../srt/mem_cache/unified_radix_cache.py | 27 +-- .../unit/mem_cache/test_swa_unittest.py | 5 +- 5 files changed, 185 insertions(+), 104 deletions(-) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index f9ecc0e69..878e530ac 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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) diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index 3ad6421a1..ee098ac95 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -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]) diff --git a/python/sglang/srt/mem_cache/swa_radix_cache.py b/python/sglang/srt/mem_cache/swa_radix_cache.py index 6e1b1c847..47682477d 100644 --- a/python/sglang/srt/mem_cache/swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/swa_radix_cache.py @@ -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 diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 608583cd7..3742905c8 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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 diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 009f956df..0e2335d0f 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -478,8 +478,9 @@ class TestSWA(unittest.TestCase): ) self.assertEqual(len(kv_indices), 6) self.assertEqual(len(last_node.key), 2) - self.assertEqual(last_node.key.token_ids[0], (5, 60)) - self.assertEqual(last_node.key.token_ids[1], (60, 70)) + # Bigram view: token_ids holds raw tokens; iteration yields bigram tuples. + self.assertTrue(last_node.key.is_bigram) + self.assertEqual(list(last_node.key), [(5, 60), (60, 70)]) def test_swa_cache_finished_req_eagle_uses_cache_protected_len_and_bigram_key(self): tree, allocator, req_to_token_pool = self._build_swa_tree(is_eagle=True)