[Perf] Make EAGLE bigram key an O(1) view on RadixKey (#23106)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user