[HiCache][SPEC] fix: normalize storage prefetch key (#23631)
This commit is contained in:
@@ -53,7 +53,6 @@ from sglang.srt.mem_cache.radix_cache import (
|
||||
compute_node_hash_values,
|
||||
split_node_hash_value,
|
||||
)
|
||||
from sglang.srt.mem_cache.utils import convert_to_bigram_key
|
||||
from sglang.srt.observability.metrics_collector import StorageMetricsCollector
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -1152,7 +1151,7 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
# todo: more policies for prefetch progress such as timeout
|
||||
# the current policy is to prefetch with best effort and terminate when queuing is over
|
||||
last_host_node, token_ids, host_indices, operation = self.ongoing_prefetch[
|
||||
last_host_node, prefetch_key, host_indices, operation = self.ongoing_prefetch[
|
||||
req_id
|
||||
]
|
||||
|
||||
@@ -1175,13 +1174,11 @@ class HiRadixCache(RadixCache):
|
||||
completed_tokens_tensor, torch.distributed.ReduceOp.MIN
|
||||
)
|
||||
min_completed_tokens = completed_tokens_tensor.item()
|
||||
fetched_token_ids = token_ids[:min_completed_tokens]
|
||||
fetched_key = prefetch_key[:min_completed_tokens]
|
||||
written_indices = host_indices[:min_completed_tokens]
|
||||
matched_length = self._insert_helper_host(
|
||||
last_host_node,
|
||||
RadixKey(
|
||||
token_ids=fetched_token_ids, extra_key=last_host_node.key.extra_key
|
||||
),
|
||||
fetched_key,
|
||||
written_indices,
|
||||
hash_value[: min_completed_tokens // self.page_size],
|
||||
)
|
||||
@@ -1192,7 +1189,7 @@ class HiRadixCache(RadixCache):
|
||||
)
|
||||
last_host_node.release_host()
|
||||
del self.ongoing_prefetch[req_id]
|
||||
self.cache_controller.prefetch_tokens_occupied -= len(token_ids)
|
||||
self.cache_controller.prefetch_tokens_occupied -= len(prefetch_key)
|
||||
|
||||
# Track tokens actually loaded from storage for this request (L3 hits)
|
||||
loaded_from_storage = min_completed_tokens - matched_length
|
||||
@@ -1269,16 +1266,14 @@ class HiRadixCache(RadixCache):
|
||||
last_hash: Optional[str] = None,
|
||||
prefix_keys: Optional[List[str]] = None,
|
||||
):
|
||||
new_input_tokens = (
|
||||
convert_to_bigram_key(new_input_tokens)
|
||||
if self.is_eagle
|
||||
else new_input_tokens
|
||||
prefetch_key = RadixKey(
|
||||
new_input_tokens,
|
||||
extra_key=last_host_node.key.extra_key,
|
||||
is_bigram=self.is_eagle,
|
||||
)
|
||||
# align the number of fetching tokens to the page size
|
||||
prefetch_length = len(new_input_tokens) - (
|
||||
len(new_input_tokens) % self.page_size
|
||||
)
|
||||
new_input_tokens = new_input_tokens[:prefetch_length]
|
||||
prefetch_key = prefetch_key.page_aligned(self.page_size)
|
||||
prefetch_length = len(prefetch_key)
|
||||
if (
|
||||
not self.enable_storage
|
||||
or prefetch_length < self.prefetch_threshold
|
||||
@@ -1306,18 +1301,18 @@ class HiRadixCache(RadixCache):
|
||||
operation = self.cache_controller.prefetch(
|
||||
req_id,
|
||||
host_indices,
|
||||
new_input_tokens,
|
||||
prefetch_key,
|
||||
last_hash,
|
||||
prefix_keys,
|
||||
**self._get_extra_pools(),
|
||||
)
|
||||
self.ongoing_prefetch[req_id] = (
|
||||
last_host_node,
|
||||
new_input_tokens,
|
||||
prefetch_key,
|
||||
host_indices,
|
||||
operation,
|
||||
)
|
||||
self.cache_controller.prefetch_tokens_occupied += len(new_input_tokens)
|
||||
self.cache_controller.prefetch_tokens_occupied += len(prefetch_key)
|
||||
|
||||
def _insert_helper_host(
|
||||
self, node: TreeNode, key: RadixKey, host_value, hash_value
|
||||
@@ -1507,7 +1502,9 @@ class HiRadixCache(RadixCache):
|
||||
if rid not in self.ongoing_prefetch:
|
||||
return
|
||||
|
||||
last_host_node, token_ids, host_indices, operation = self.ongoing_prefetch[rid]
|
||||
last_host_node, prefetch_key, host_indices, operation = self.ongoing_prefetch[
|
||||
rid
|
||||
]
|
||||
if operation.host_indices is None:
|
||||
return
|
||||
|
||||
@@ -1516,4 +1513,4 @@ class HiRadixCache(RadixCache):
|
||||
last_host_node.release_host()
|
||||
del self.ongoing_prefetch[rid]
|
||||
self.cache_controller.append_host_mem_release(host_indices[:completed_tokens])
|
||||
self.cache_controller.prefetch_tokens_occupied -= len(token_ids)
|
||||
self.cache_controller.prefetch_tokens_occupied -= len(prefetch_key)
|
||||
|
||||
Reference in New Issue
Block a user