[Refactor] Clean up radix cache related API (#7303)
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
@@ -38,7 +38,7 @@ import logging
|
||||
import threading
|
||||
from enum import Enum, auto
|
||||
from http import HTTPStatus
|
||||
from typing import TYPE_CHECKING, List, Optional, Set, Tuple, Union
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Set, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -436,7 +436,7 @@ class Req:
|
||||
self,
|
||||
rid: str,
|
||||
origin_input_text: str,
|
||||
origin_input_ids: Tuple[int],
|
||||
origin_input_ids: List[int],
|
||||
sampling_params: SamplingParams,
|
||||
return_logprob: bool = False,
|
||||
top_logprobs_num: int = 0,
|
||||
@@ -467,7 +467,7 @@ class Req:
|
||||
# Each decode stage's output ids
|
||||
self.output_ids = []
|
||||
# fill_ids = origin_input_ids + output_ids. Updated if chunked.
|
||||
self.fill_ids = None
|
||||
self.fill_ids = []
|
||||
self.session_id = session_id
|
||||
self.input_embeds = input_embeds
|
||||
|
||||
@@ -519,13 +519,14 @@ class Req:
|
||||
|
||||
# Prefix info
|
||||
# The indices to kv cache for the shared prefix.
|
||||
self.prefix_indices = []
|
||||
self.prefix_indices: torch.Tensor = []
|
||||
# Number of tokens to run prefill.
|
||||
self.extend_input_len = 0
|
||||
# The relative logprob_start_len in an extend batch
|
||||
self.extend_logprob_start_len = 0
|
||||
self.last_node = None
|
||||
self.last_node_global = None
|
||||
self.last_node: Any = None
|
||||
self.last_host_node: Any = None
|
||||
self.host_hit_length = 0
|
||||
|
||||
# Whether or not if it is chunked. It increments whenever
|
||||
# it is chunked, and decrement whenever chunked request is
|
||||
@@ -644,21 +645,17 @@ class Req:
|
||||
def init_next_round_input(
|
||||
self,
|
||||
tree_cache: Optional[BasePrefixCache] = None,
|
||||
enable_hierarchical_cache=False,
|
||||
):
|
||||
self.fill_ids = self.origin_input_ids + self.output_ids
|
||||
if tree_cache is not None:
|
||||
# tree cache is None if the prefix is not computed with tree cache.
|
||||
if enable_hierarchical_cache:
|
||||
self.prefix_indices, self.last_node, self.last_node_global = (
|
||||
tree_cache.match_prefix(
|
||||
key=self.adjust_max_prefix_ids(), include_evicted=True
|
||||
)
|
||||
)
|
||||
else:
|
||||
self.prefix_indices, self.last_node = tree_cache.match_prefix(
|
||||
rid=self.rid, key=self.adjust_max_prefix_ids()
|
||||
)
|
||||
(
|
||||
self.prefix_indices,
|
||||
self.last_node,
|
||||
self.last_host_node,
|
||||
self.host_hit_length,
|
||||
) = tree_cache.match_prefix(
|
||||
key=self.adjust_max_prefix_ids(),
|
||||
)
|
||||
self.extend_input_len = len(self.fill_ids) - len(self.prefix_indices)
|
||||
|
||||
def adjust_max_prefix_ids(self):
|
||||
|
||||
@@ -90,7 +90,7 @@ class SchedulePolicy:
|
||||
def calc_priority(self, waiting_queue: List[Req]) -> bool:
|
||||
if self.policy == CacheAgnosticPolicy.FCFS:
|
||||
# A shortcut for FCFS
|
||||
return
|
||||
return False
|
||||
|
||||
policy = self._determine_active_policy(waiting_queue)
|
||||
|
||||
@@ -134,7 +134,7 @@ class SchedulePolicy:
|
||||
"""
|
||||
try:
|
||||
policy_enum = CacheAwarePolicy(policy)
|
||||
if tree_cache.disable:
|
||||
if getattr(tree_cache, "disable", True):
|
||||
# If tree_cache is disabled, using CacheAgnosticPolicy policy
|
||||
return CacheAgnosticPolicy.FCFS
|
||||
return policy_enum
|
||||
@@ -158,14 +158,9 @@ class SchedulePolicy:
|
||||
prefix_ids = r.adjust_max_prefix_ids()
|
||||
|
||||
# NOTE: the prefix_indices must always be aligned with last_node
|
||||
if self.enable_hierarchical_cache:
|
||||
r.prefix_indices, r.last_node, r.last_node_global = (
|
||||
self.tree_cache.match_prefix(key=prefix_ids, include_evicted=True)
|
||||
)
|
||||
else:
|
||||
r.prefix_indices, r.last_node = self.tree_cache.match_prefix(
|
||||
rid=r.rid, key=prefix_ids
|
||||
)
|
||||
r.prefix_indices, r.last_node, r.last_host_node, r.host_hit_length = (
|
||||
self.tree_cache.match_prefix(rid=r.rid, key=prefix_ids)
|
||||
)
|
||||
|
||||
# NOTE(sang): This logic is for in-batch prefix caching;
|
||||
# If there are more than 1 request that have small matching prefix from
|
||||
@@ -175,7 +170,7 @@ class SchedulePolicy:
|
||||
# threshold means we cannot use in-batch prefix caching for short prefixes.
|
||||
# It is kind of common when the engine is long running (e.g., imagine the prefix "the").
|
||||
if len(r.prefix_indices) <= IN_BATCH_PREFIX_CACHING_CHECK_THRESHOLD:
|
||||
in_batch_matching_prefixes, _ = (
|
||||
in_batch_matching_prefixes, _, _, _ = (
|
||||
self.waiting_queue_radix_tree.match_prefix(
|
||||
rid=r.rid, key=prefix_ids
|
||||
)
|
||||
@@ -268,6 +263,7 @@ class AddReqResult(Enum):
|
||||
class PrefillAdder:
|
||||
def __init__(
|
||||
self,
|
||||
page_size: int,
|
||||
tree_cache: BasePrefixCache,
|
||||
token_to_kv_pool_allocator: TokenToKVPoolAllocator,
|
||||
running_batch: ScheduleBatch,
|
||||
@@ -276,6 +272,7 @@ class PrefillAdder:
|
||||
rem_chunk_tokens: Optional[int],
|
||||
mixed_with_decode_tokens: int = 0,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self.tree_cache = tree_cache
|
||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||
self.running_batch = running_batch
|
||||
@@ -442,46 +439,43 @@ class PrefillAdder:
|
||||
|
||||
return self.budget_state()
|
||||
|
||||
def add_one_req(
|
||||
self, req: Req, has_chunked_req: bool, enable_hierarchical_cache: bool = False
|
||||
):
|
||||
def add_one_req(self, req: Req, has_chunked_req: bool):
|
||||
if req.sampling_params.ignore_eos and getattr(self.tree_cache, "disable", True):
|
||||
return self.add_one_req_ignore_eos(req, has_chunked_req)
|
||||
|
||||
total_tokens = req.extend_input_len + min(
|
||||
req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS_ESTIMATION
|
||||
)
|
||||
input_tokens = (
|
||||
-(-req.extend_input_len // self.tree_cache.page_size)
|
||||
* self.tree_cache.page_size
|
||||
)
|
||||
|
||||
# adjusting the input_tokens based on host_hit_length and page_size
|
||||
real_input_tokens = req.extend_input_len - req.host_hit_length
|
||||
real_input_tokens = -(-real_input_tokens // self.page_size) * self.page_size
|
||||
prefix_len = len(req.prefix_indices)
|
||||
|
||||
if total_tokens >= self.rem_total_tokens:
|
||||
return AddReqResult.NO_TOKEN
|
||||
|
||||
if input_tokens > self.rem_input_tokens and len(self.can_run_list) != 0:
|
||||
if real_input_tokens >= self.rem_input_tokens and len(self.can_run_list) != 0:
|
||||
return AddReqResult.OTHER
|
||||
|
||||
with self._lock_node(req.last_node):
|
||||
if total_tokens > self.rem_total_tokens:
|
||||
# self.rem_total_tokens may decrease after the lock acquisition
|
||||
if total_tokens >= self.rem_total_tokens:
|
||||
return AddReqResult.NO_TOKEN
|
||||
|
||||
if (
|
||||
enable_hierarchical_cache
|
||||
and req.last_node_global is not None
|
||||
and req.last_node_global.evicted
|
||||
):
|
||||
req.last_node, req.prefix_indices = self.tree_cache.init_load_back(
|
||||
req.last_node_global, req.prefix_indices
|
||||
if req.host_hit_length > 0:
|
||||
new_indices, req.last_node = self.tree_cache.init_load_back(
|
||||
req.last_host_node, req.host_hit_length
|
||||
)
|
||||
req.prefix_indices = torch.cat([req.prefix_indices, new_indices])
|
||||
req.extend_input_len = len(req.fill_ids) - len(req.prefix_indices)
|
||||
input_tokens = (
|
||||
-(-req.extend_input_len // self.tree_cache.page_size)
|
||||
* self.tree_cache.page_size
|
||||
)
|
||||
prefix_len = len(req.prefix_indices)
|
||||
|
||||
input_tokens = -(-req.extend_input_len // self.page_size) * self.page_size
|
||||
|
||||
if input_tokens >= self.rem_input_tokens and len(self.can_run_list) != 0:
|
||||
return AddReqResult.OTHER
|
||||
|
||||
if self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
||||
# Non-chunked prefill
|
||||
self.can_run_list.append(req)
|
||||
@@ -496,7 +490,7 @@ class PrefillAdder:
|
||||
)
|
||||
else:
|
||||
# Make sure at least one page is available
|
||||
trunc_len = self.rem_chunk_tokens - self.tree_cache.page_size + 1
|
||||
trunc_len = self.rem_chunk_tokens - self.page_size + 1
|
||||
if trunc_len <= 0:
|
||||
return AddReqResult.OTHER
|
||||
|
||||
|
||||
@@ -1467,15 +1467,14 @@ class Scheduler(
|
||||
return None
|
||||
|
||||
if self.enable_hierarchical_cache:
|
||||
# check for completion of hierarchical cache activities to release memory
|
||||
self.tree_cache.writing_check()
|
||||
self.tree_cache.loading_check()
|
||||
self.tree_cache.check_hicache_events()
|
||||
|
||||
# Get priority queue
|
||||
prefix_computed = self.policy.calc_priority(self.waiting_queue)
|
||||
self.policy.calc_priority(self.waiting_queue)
|
||||
|
||||
# Prefill policy
|
||||
adder = PrefillAdder(
|
||||
self.page_size,
|
||||
self.tree_cache,
|
||||
self.token_to_kv_pool_allocator,
|
||||
self.running_batch,
|
||||
@@ -1517,19 +1516,8 @@ class Scheduler(
|
||||
self.running_batch.batch_is_full = True
|
||||
break
|
||||
|
||||
# bypass prefix_computed if enable_hierarchical_cache
|
||||
req.init_next_round_input(
|
||||
(
|
||||
None
|
||||
if (prefix_computed and not self.enable_hierarchical_cache)
|
||||
else self.tree_cache
|
||||
),
|
||||
self.enable_hierarchical_cache,
|
||||
)
|
||||
|
||||
res = adder.add_one_req(
|
||||
req, self.chunked_req, self.enable_hierarchical_cache
|
||||
)
|
||||
req.init_next_round_input(self.tree_cache)
|
||||
res = adder.add_one_req(req, has_chunked_req=(self.chunked_req is not None))
|
||||
|
||||
if res != AddReqResult.CONTINUE:
|
||||
if res == AddReqResult.NO_TOKEN:
|
||||
@@ -1581,7 +1569,9 @@ class Scheduler(
|
||||
)
|
||||
if self.enable_hierarchical_cache:
|
||||
# todo (zhiqiang): disable cuda graph execution if hicache loading triggered
|
||||
new_batch.hicache_consumer_index = self.tree_cache.ready_to_load_cache()
|
||||
new_batch.hicache_consumer_index = (
|
||||
self.tree_cache.ready_to_load_host_cache()
|
||||
)
|
||||
|
||||
new_batch.prepare_for_extend()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user