[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):
|
def match_prefix(self, params: MatchPrefixParams):
|
||||||
key = params.key
|
key = params.key
|
||||||
empty_value = torch.empty((0,), dtype=torch.int64, device=self.device)
|
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:
|
if self.disable or len(key) == 0:
|
||||||
return MatchResult(
|
return MatchResult(
|
||||||
device_indices=empty_value,
|
device_indices=empty_value,
|
||||||
@@ -1394,7 +1394,7 @@ class HiRadixCache(RadixCache):
|
|||||||
|
|
||||||
if priority is None:
|
if priority is None:
|
||||||
priority = 0
|
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:
|
if len(key) == 0:
|
||||||
return InsertResult(prefix_len=0)
|
return InsertResult(prefix_len=0)
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
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
|
Copyright 2023-2024 SGLang Team
|
||||||
@@ -62,58 +61,95 @@ from sglang.srt.mem_cache.evict_policy import (
|
|||||||
PriorityStrategy,
|
PriorityStrategy,
|
||||||
SLRUStrategy,
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
|
||||||
|
|
||||||
class RadixKey:
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
token_ids: List[int],
|
token_ids: List[int],
|
||||||
extra_key: Optional[str] = None,
|
extra_key: Optional[str] = None,
|
||||||
is_bigram: bool = False,
|
is_bigram: bool = False,
|
||||||
):
|
):
|
||||||
# token ids sequence
|
# token ids sequence (raw ints in both modes)
|
||||||
self.token_ids = token_ids
|
self.token_ids = token_ids
|
||||||
# extra key (e.g. lora_id, cache_salt)
|
# extra key (e.g. lora_id, cache_salt)
|
||||||
self.extra_key = extra_key
|
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
|
self.is_bigram = is_bigram
|
||||||
|
|
||||||
def __len__(self) -> int:
|
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)
|
return len(self.token_ids)
|
||||||
|
|
||||||
def __iter__(self) -> Iterator[int]:
|
def __iter__(self) -> Iterator:
|
||||||
return iter(self.token_ids)
|
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":
|
def __getitem__(self, idx: Union[int, slice]) -> "RadixKey":
|
||||||
if isinstance(idx, slice):
|
# Normalize int -> 1-element slice so the rest handles one shape.
|
||||||
return RadixKey(self.token_ids[idx], self.extra_key)
|
if isinstance(idx, int):
|
||||||
return RadixKey([self.token_ids[idx]], self.extra_key)
|
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:
|
def __repr__(self) -> str:
|
||||||
preview = self.token_ids[:10]
|
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(
|
||||||
def maybe_bigram_convert(
|
self,
|
||||||
is_eagle: bool,
|
is_eagle: bool,
|
||||||
key: RadixKey,
|
|
||||||
value: Optional[torch.Tensor] = None,
|
value: Optional[torch.Tensor] = None,
|
||||||
) -> Tuple[RadixKey, Optional[torch.Tensor]]:
|
) -> Tuple["RadixKey", Optional[torch.Tensor]]:
|
||||||
if is_eagle and not key.is_bigram:
|
# O(1): flip the bigram flag instead of materializing a tuple list.
|
||||||
key.token_ids = convert_to_bigram_key(key.token_ids)
|
# value is paired with raw tokens and gets truncated to the bigram count.
|
||||||
key.is_bigram = True
|
if is_eagle and not self.is_bigram:
|
||||||
|
self.is_bigram = True
|
||||||
if value is not None:
|
if value is not None:
|
||||||
value = value[: len(key)]
|
value = value[: len(self)]
|
||||||
return key, value
|
return self, value
|
||||||
|
|
||||||
|
|
||||||
def page_align_keys(key: list, page_size) -> list:
|
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.
|
||||||
|
|
||||||
|
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:
|
if page_size == 1:
|
||||||
return key
|
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
|
page_aligned_len = len(key) // page_size * page_size
|
||||||
return key[:page_aligned_len]
|
return key[:page_aligned_len]
|
||||||
|
|
||||||
@@ -190,28 +226,56 @@ def _check_extra_key(key0: RadixKey, key1: RadixKey):
|
|||||||
|
|
||||||
def _key_match_page_size1(key0: RadixKey, key1: RadixKey):
|
def _key_match_page_size1(key0: RadixKey, key1: RadixKey):
|
||||||
_check_extra_key(key0, key1)
|
_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
|
i = 0
|
||||||
for k0, k1 in zip(key0.token_ids, key1.token_ids):
|
for a, b in zip(t0, t1):
|
||||||
if k0 != k1:
|
if a != b:
|
||||||
break
|
break
|
||||||
i += 1
|
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
|
return i
|
||||||
|
|
||||||
|
|
||||||
def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
|
def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
|
||||||
_check_extra_key(key0, key1)
|
_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
|
i = 0
|
||||||
while i < min_len:
|
while i < min_len:
|
||||||
if key0.token_ids[i : i + page_size] != key1.token_ids[i : i + page_size]:
|
if key0.token_ids[i : i + page_size] != key1.token_ids[i : i + page_size]:
|
||||||
break
|
break
|
||||||
i += page_size
|
i += page_size
|
||||||
|
|
||||||
return i
|
return i
|
||||||
|
|
||||||
|
|
||||||
def get_child_key(key: RadixKey, page_size: int = 1):
|
def get_child_key(key: RadixKey, page_size: int = 1):
|
||||||
|
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:
|
||||||
if page_size == 1:
|
if page_size == 1:
|
||||||
plain_key = key.token_ids[0]
|
plain_key = key.token_ids[0]
|
||||||
else:
|
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]:
|
def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]:
|
||||||
"""Compute SHA256-based hash values for position-aware identification.
|
"""Compute SHA256-based hash values for position-aware identification.
|
||||||
|
|
||||||
Args:
|
In bigram mode, each page logically covers `page_size` bigrams over
|
||||||
node: The TreeNode to compute hash values for
|
`page_size + 1` raw tokens; we feed overlapping (t_i, t_{i+1}) byte pairs
|
||||||
page_size: The page size for chunking tokens
|
to the hasher so the output matches the pre-optimization tuple-based hash.
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of SHA256 hex strings, one per page
|
|
||||||
"""
|
"""
|
||||||
hash_values = []
|
hash_values = []
|
||||||
|
|
||||||
# Get parent's last hash value if parent exists
|
|
||||||
parent_hash = None
|
parent_hash = None
|
||||||
if node.parent is not None and node.parent.hash_value is not 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:
|
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
|
||||||
parent_hash = node.parent.hash_value[-1]
|
parent_hash = node.parent.hash_value[-1]
|
||||||
|
|
||||||
# Iterate through node's pages
|
raw = node.key.token_ids
|
||||||
for start in range(0, len(node.key), page_size):
|
is_bigram = node.key.is_bigram
|
||||||
page_tokens = node.key.token_ids[start : start + page_size]
|
logical_len = len(node.key)
|
||||||
if not page_tokens:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Use SHA256-based chaining via get_hash_str
|
for start in range(0, logical_len, page_size):
|
||||||
hash_val = get_hash_str(page_tokens, prior_hash=parent_hash)
|
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)
|
hash_values.append(hash_val)
|
||||||
parent_hash = hash_val
|
parent_hash = hash_val
|
||||||
|
|
||||||
return hash_values
|
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(
|
def split_node_hash_value(
|
||||||
child_hash_value: Optional[List[str]], split_len: int, page_size: int
|
child_hash_value: Optional[List[str]], split_len: int, page_size: int
|
||||||
) -> tuple[Optional[List[str]], Optional[List[str]]]:
|
) -> tuple[Optional[List[str]], Optional[List[str]]]:
|
||||||
@@ -366,11 +450,6 @@ class RadixCache(BasePrefixCache):
|
|||||||
self.evictable_leaves.clear()
|
self.evictable_leaves.clear()
|
||||||
self._record_all_cleared_event()
|
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:
|
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
|
||||||
"""Find the longest cached prefix of ``key`` in the radix tree.
|
"""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.
|
subsequent match efficiency and does not duplicate data.
|
||||||
"""
|
"""
|
||||||
key = params.key
|
key = params.key
|
||||||
key, _ = self.maybe_bigram_convert(key)
|
key, _ = key.maybe_to_bigram_view(self.is_eagle)
|
||||||
|
|
||||||
def empty_match_result():
|
def empty_match_result():
|
||||||
return MatchResult(
|
return MatchResult(
|
||||||
@@ -453,9 +532,11 @@ class RadixCache(BasePrefixCache):
|
|||||||
chunked = params.chunked
|
chunked = params.chunked
|
||||||
|
|
||||||
if value is None:
|
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)
|
prefix_len = self._insert_helper(self.root_node, key, value, priority, chunked)
|
||||||
return InsertResult(prefix_len=prefix_len)
|
return InsertResult(prefix_len=prefix_len)
|
||||||
@@ -479,11 +560,9 @@ class RadixCache(BasePrefixCache):
|
|||||||
req.req_pool_idx, : len(token_ids)
|
req.req_pool_idx, : len(token_ids)
|
||||||
]
|
]
|
||||||
|
|
||||||
# Maybe convert to bigram keys for EAGLE
|
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_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)
|
|
||||||
radix_key = RadixKey(keys, req.extra_key, 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
|
# Radix Cache takes one ref in memory pool
|
||||||
if is_insert:
|
if is_insert:
|
||||||
@@ -498,11 +577,11 @@ class RadixCache(BasePrefixCache):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.token_to_kv_pool_allocator.free(
|
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
|
# 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
|
# Remove req slot release the cache lock
|
||||||
self.dec_lock_ref(req.last_node)
|
self.dec_lock_ref(req.last_node)
|
||||||
@@ -517,11 +596,9 @@ class RadixCache(BasePrefixCache):
|
|||||||
req.req_pool_idx, : len(token_ids)
|
req.req_pool_idx, : len(token_ids)
|
||||||
]
|
]
|
||||||
|
|
||||||
# Maybe convert to bigram keys for EAGLE
|
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_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)
|
|
||||||
radix_key = RadixKey(keys, req.extra_key, 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
|
# Radix Cache takes one ref in memory pool
|
||||||
result = self.insert(
|
result = self.insert(
|
||||||
@@ -544,7 +621,9 @@ class RadixCache(BasePrefixCache):
|
|||||||
match_result.device_indices,
|
match_result.device_indices,
|
||||||
match_result.last_device_node,
|
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(
|
self.req_to_token_pool.write(
|
||||||
(req.req_pool_idx, slice(req.cache_protected_len, len(new_indices))),
|
(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])
|
parent_block_hash = hash_str_to_int64(node.parent.hash_value[-1])
|
||||||
|
|
||||||
page_index = 0
|
page_index = 0
|
||||||
for start in range(0, len(node.key), self.page_size):
|
logical_len = len(node.key)
|
||||||
page_tokens = node.key.token_ids[start : start + self.page_size]
|
is_bigram = node.key.is_bigram
|
||||||
if not page_tokens:
|
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
|
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])
|
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)
|
node.hash_value = compute_node_hash_values(node, self.page_size)
|
||||||
|
|
||||||
page_index = 0
|
page_index = 0
|
||||||
for start in range(0, len(node.key), self.page_size):
|
logical_len = len(node.key)
|
||||||
page_tokens = node.key.token_ids[start : start + self.page_size]
|
for start in range(0, logical_len, self.page_size):
|
||||||
if not page_tokens:
|
end = min(start + self.page_size, logical_len)
|
||||||
|
if end <= start:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
block_hash = hash_str_to_int64(node.hash_value[page_index])
|
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_page_size1,
|
||||||
_key_match_paged,
|
_key_match_paged,
|
||||||
get_child_key,
|
get_child_key,
|
||||||
maybe_bigram_convert,
|
|
||||||
page_align_keys,
|
page_align_keys,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
@@ -432,9 +431,9 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
swa_evicted_seqlen = params.swa_evicted_seqlen
|
swa_evicted_seqlen = params.swa_evicted_seqlen
|
||||||
|
|
||||||
if value is None:
|
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(
|
prefix_len = self._insert_helper(
|
||||||
self.root_node, key, value, prev_prefix_len, swa_evicted_seqlen
|
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
|
req.req_pool_idx, :kv_committed_len
|
||||||
]
|
]
|
||||||
|
|
||||||
# Maybe convert to bigram keys for EAGLE
|
# EAGLE: skip tuple materialization; is_bigram flag gives bigram semantics.
|
||||||
keys = self.key_convert_fn(token_ids)
|
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
|
||||||
keys = page_align_keys(keys, self.page_size)
|
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
|
||||||
page_aligned_len = len(keys)
|
page_aligned_len = len(radix_key)
|
||||||
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
|
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
|
old_prefix_len = req.cache_protected_len
|
||||||
|
|
||||||
# Radix Cache takes one ref in memory pool
|
# Radix Cache takes one ref in memory pool
|
||||||
@@ -508,10 +502,9 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
req.req_pool_idx, : len(token_ids)
|
req.req_pool_idx, : len(token_ids)
|
||||||
]
|
]
|
||||||
|
|
||||||
keys = self.key_convert_fn(token_ids)
|
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
|
||||||
keys = page_align_keys(keys, self.page_size)
|
|
||||||
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
|
|
||||||
radix_key = RadixKey(keys, req.extra_key, 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
|
old_prefix_len = req.cache_protected_len
|
||||||
|
|
||||||
# Radix Cache takes one ref in memory pool
|
# Radix Cache takes one ref in memory pool
|
||||||
@@ -846,7 +839,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
def _match_pre_processor(self, params: MatchPrefixParams) -> Optional[RadixKey]:
|
def _match_pre_processor(self, params: MatchPrefixParams) -> Optional[RadixKey]:
|
||||||
"""Preprocess the key before matching."""
|
"""Preprocess the key before matching."""
|
||||||
key = params.key
|
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:
|
if self.disable or len(key) == 0:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -25,7 +25,6 @@ from sglang.srt.mem_cache.radix_cache import (
|
|||||||
_key_match_page_size1,
|
_key_match_page_size1,
|
||||||
_key_match_paged,
|
_key_match_paged,
|
||||||
get_child_key,
|
get_child_key,
|
||||||
maybe_bigram_convert,
|
|
||||||
page_align_keys,
|
page_align_keys,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_cache_components import (
|
from sglang.srt.mem_cache.unified_cache_components import (
|
||||||
@@ -239,7 +238,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
key = params.key
|
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:
|
if self.disable or len(key) == 0:
|
||||||
return MatchResult(
|
return MatchResult(
|
||||||
device_indices=torch.empty(
|
device_indices=torch.empty(
|
||||||
@@ -264,9 +263,9 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
key = params.key
|
key = params.key
|
||||||
value = params.value
|
value = params.value
|
||||||
if value is None:
|
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)
|
result = self._insert_helper(self.root_node, key, value, params)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@@ -355,12 +354,11 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
token_ids = token_ids[:effective_cache_len]
|
token_ids = token_ids[:effective_cache_len]
|
||||||
kv_indices = kv_indices[:effective_cache_len]
|
kv_indices = kv_indices[:effective_cache_len]
|
||||||
|
|
||||||
# Key convert + page align
|
# Page align on raw tokens; bigram semantics via is_bigram flag.
|
||||||
keys = self.key_convert_fn(token_ids)
|
keys = page_align_keys(token_ids, self.page_size, is_bigram=self.is_eagle)
|
||||||
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)
|
|
||||||
radix_key = RadixKey(keys, req.extra_key, 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.key = radix_key
|
||||||
insert_params.value = values
|
insert_params.value = values
|
||||||
@@ -422,12 +420,13 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
|
|
||||||
kv_indices = kv_indices_orig[:effective_cache_len]
|
kv_indices = kv_indices_orig[:effective_cache_len]
|
||||||
|
|
||||||
# Key convert + page align
|
# Page align on raw tokens; bigram semantics via is_bigram flag.
|
||||||
keys = self.key_convert_fn(token_ids[:effective_cache_len])
|
keys = page_align_keys(
|
||||||
keys = page_align_keys(keys, self.page_size)
|
token_ids[:effective_cache_len], self.page_size, is_bigram=self.is_eagle
|
||||||
page_aligned_len = len(keys)
|
)
|
||||||
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
|
|
||||||
radix_key = RadixKey(keys, req.extra_key, 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.key = radix_key
|
||||||
insert_params.value = values
|
insert_params.value = values
|
||||||
|
|||||||
@@ -478,8 +478,9 @@ class TestSWA(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
self.assertEqual(len(kv_indices), 6)
|
self.assertEqual(len(kv_indices), 6)
|
||||||
self.assertEqual(len(last_node.key), 2)
|
self.assertEqual(len(last_node.key), 2)
|
||||||
self.assertEqual(last_node.key.token_ids[0], (5, 60))
|
# Bigram view: token_ids holds raw tokens; iteration yields bigram tuples.
|
||||||
self.assertEqual(last_node.key.token_ids[1], (60, 70))
|
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):
|
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)
|
tree, allocator, req_to_token_pool = self._build_swa_tree(is_eagle=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user