[Refactor] Replace page_align_keys helper with RadixKey.page_aligned method (#23107)
This commit is contained in:
@@ -1216,10 +1216,8 @@ class HiRadixCache(RadixCache):
|
||||
host_hit_length=0,
|
||||
)
|
||||
|
||||
key = key.page_aligned(self.page_size)
|
||||
page_aligned_len = len(key)
|
||||
if self.page_size != 1:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
|
||||
value, last_node = self._match_prefix_helper(self.root_node, key)
|
||||
if value:
|
||||
@@ -1394,15 +1392,15 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
if priority is None:
|
||||
priority = 0
|
||||
|
||||
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
|
||||
key = key.page_aligned(self.page_size)
|
||||
if value is not None:
|
||||
value = value[: len(key)]
|
||||
|
||||
if len(key) == 0:
|
||||
return InsertResult(prefix_len=0)
|
||||
|
||||
if self.is_eagle and value is not None:
|
||||
# Make sure the value len equal to the EAGLE bigram key len
|
||||
value = value[: len(key)]
|
||||
|
||||
node = self.root_node
|
||||
child_key = self.get_child_key_fn(key)
|
||||
total_prefix_length = 0
|
||||
|
||||
@@ -122,6 +122,12 @@ class RadixKey:
|
||||
preview = self.token_ids[:10]
|
||||
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 page_aligned(self, page_size: int) -> "RadixKey":
|
||||
if page_size == 1:
|
||||
return self
|
||||
aligned_len = len(self) // page_size * page_size
|
||||
return self[:aligned_len]
|
||||
|
||||
def maybe_to_bigram_view(
|
||||
self,
|
||||
is_eagle: bool,
|
||||
@@ -136,24 +142,6 @@ class RadixKey:
|
||||
return self, 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.
|
||||
|
||||
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]
|
||||
|
||||
|
||||
class TreeNode:
|
||||
|
||||
counter = 0
|
||||
@@ -504,9 +492,7 @@ class RadixCache(BasePrefixCache):
|
||||
if self.disable or len(key) == 0:
|
||||
return empty_match_result()
|
||||
|
||||
if self.page_size != 1:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
key = key.page_aligned(self.page_size)
|
||||
|
||||
if len(key) == 0:
|
||||
return empty_match_result()
|
||||
@@ -531,12 +517,13 @@ class RadixCache(BasePrefixCache):
|
||||
priority = params.priority
|
||||
chunked = params.chunked
|
||||
|
||||
if value is None:
|
||||
# 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 = key.maybe_to_bigram_view(self.is_eagle, value)
|
||||
key = key.page_aligned(self.page_size)
|
||||
if value is not None:
|
||||
value = value[: len(key)]
|
||||
else:
|
||||
# Debug/test fallback: use token ids themselves as values.
|
||||
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
|
||||
|
||||
prefix_len = self._insert_helper(self.root_node, key, value, priority, chunked)
|
||||
return InsertResult(prefix_len=prefix_len)
|
||||
@@ -560,9 +547,11 @@ class RadixCache(BasePrefixCache):
|
||||
req.req_pool_idx, : len(token_ids)
|
||||
]
|
||||
|
||||
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_key = RadixKey(
|
||||
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||
).page_aligned(self.page_size)
|
||||
key_len = len(radix_key)
|
||||
values = kv_indices[:key_len].to(dtype=torch.int64, copy=True)
|
||||
|
||||
# Radix Cache takes one ref in memory pool
|
||||
if is_insert:
|
||||
@@ -577,11 +566,11 @@ class RadixCache(BasePrefixCache):
|
||||
)
|
||||
else:
|
||||
self.token_to_kv_pool_allocator.free(
|
||||
kv_indices[req.cache_protected_len : len(radix_key)]
|
||||
kv_indices[req.cache_protected_len : key_len]
|
||||
)
|
||||
|
||||
# free the unaligned tail
|
||||
self.token_to_kv_pool_allocator.free(kv_indices[len(radix_key) :])
|
||||
self.token_to_kv_pool_allocator.free(kv_indices[key_len:])
|
||||
|
||||
# Remove req slot release the cache lock
|
||||
self.dec_lock_ref(req.last_node)
|
||||
@@ -596,8 +585,9 @@ class RadixCache(BasePrefixCache):
|
||||
req.req_pool_idx, : len(token_ids)
|
||||
]
|
||||
|
||||
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)
|
||||
radix_key = RadixKey(
|
||||
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||
).page_aligned(self.page_size)
|
||||
values = kv_indices[: len(radix_key)].to(dtype=torch.int64, copy=True)
|
||||
|
||||
# Radix Cache takes one ref in memory pool
|
||||
|
||||
@@ -46,7 +46,6 @@ from sglang.srt.mem_cache.radix_cache import (
|
||||
_key_match_page_size1,
|
||||
_key_match_paged,
|
||||
get_child_key,
|
||||
page_align_keys,
|
||||
)
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.utils import convert_to_bigram_key
|
||||
@@ -430,10 +429,12 @@ class SWARadixCache(BasePrefixCache):
|
||||
prev_prefix_len = params.prev_prefix_len
|
||||
swa_evicted_seqlen = params.swa_evicted_seqlen
|
||||
|
||||
if value is None:
|
||||
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
|
||||
|
||||
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
|
||||
key = key.page_aligned(self.page_size)
|
||||
if value is not None:
|
||||
value = value[: len(key)]
|
||||
else:
|
||||
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
|
||||
|
||||
prefix_len = self._insert_helper(
|
||||
self.root_node, key, value, prev_prefix_len, swa_evicted_seqlen
|
||||
@@ -455,9 +456,9 @@ class SWARadixCache(BasePrefixCache):
|
||||
req.req_pool_idx, :kv_committed_len
|
||||
]
|
||||
|
||||
# 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)
|
||||
radix_key = RadixKey(
|
||||
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||
).page_aligned(self.page_size)
|
||||
page_aligned_len = len(radix_key)
|
||||
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
|
||||
old_prefix_len = req.cache_protected_len
|
||||
@@ -502,8 +503,9 @@ class SWARadixCache(BasePrefixCache):
|
||||
req.req_pool_idx, : len(token_ids)
|
||||
]
|
||||
|
||||
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)
|
||||
radix_key = RadixKey(
|
||||
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||
).page_aligned(self.page_size)
|
||||
values = kv_indices[: len(radix_key)].to(dtype=torch.int64, copy=True)
|
||||
old_prefix_len = req.cache_protected_len
|
||||
|
||||
@@ -840,14 +842,11 @@ class SWARadixCache(BasePrefixCache):
|
||||
"""Preprocess the key before matching."""
|
||||
key = params.key
|
||||
key, _ = key.maybe_to_bigram_view(self.is_eagle)
|
||||
|
||||
if self.disable or len(key) == 0:
|
||||
return None
|
||||
|
||||
if self.page_size != 1:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
|
||||
key = key.page_aligned(self.page_size)
|
||||
if len(key) == 0:
|
||||
return None
|
||||
return key
|
||||
|
||||
def _match_post_processor(
|
||||
|
||||
@@ -25,7 +25,6 @@ from sglang.srt.mem_cache.radix_cache import (
|
||||
_key_match_page_size1,
|
||||
_key_match_paged,
|
||||
get_child_key,
|
||||
page_align_keys,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache_components import (
|
||||
_NUM_COMPONENT_TYPES,
|
||||
@@ -249,9 +248,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
last_device_node=self.root_node,
|
||||
last_host_node=self.root_node,
|
||||
)
|
||||
if self.page_size != 1:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
key = key.page_aligned(self.page_size)
|
||||
|
||||
value, last_node, best_value_len = self._match_prefix_helper(key)
|
||||
return self._match_post_processor(params, value, last_node, best_value_len)
|
||||
@@ -262,10 +259,13 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
key = params.key
|
||||
value = params.value
|
||||
if value is None:
|
||||
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
|
||||
key = key.page_aligned(self.page_size)
|
||||
if value is not None:
|
||||
value = value[: len(key)]
|
||||
else:
|
||||
value = torch.tensor(key.token_ids[: len(key)], dtype=torch.int64)
|
||||
|
||||
key, value = key.maybe_to_bigram_view(self.is_eagle, value)
|
||||
result = self._insert_helper(self.root_node, key, value, params)
|
||||
return result
|
||||
|
||||
@@ -354,9 +354,9 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
token_ids = token_ids[:effective_cache_len]
|
||||
kv_indices = kv_indices[:effective_cache_len]
|
||||
|
||||
# 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)
|
||||
radix_key = RadixKey(
|
||||
token_ids, req.extra_key, is_bigram=self.is_eagle
|
||||
).page_aligned(self.page_size)
|
||||
page_aligned_len = len(radix_key)
|
||||
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
|
||||
|
||||
@@ -420,11 +420,11 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
kv_indices = kv_indices_orig[:effective_cache_len]
|
||||
|
||||
# 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)
|
||||
radix_key = RadixKey(
|
||||
token_ids[:effective_cache_len],
|
||||
req.extra_key,
|
||||
is_bigram=self.is_eagle,
|
||||
).page_aligned(self.page_size)
|
||||
page_aligned_len = len(radix_key)
|
||||
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user