Refactor kv cache free (#11351)

This commit is contained in:
cctry
2025-10-14 17:45:19 -07:00
committed by GitHub
parent 325951460f
commit 1d7f783501
8 changed files with 72 additions and 90 deletions

View File

@@ -40,7 +40,7 @@ class BasePrefixCache(ABC):
pass
@abstractmethod
def cache_finished_req(self, req: Req, **kwargs):
def cache_finished_req(self, req: Req, is_insert: bool = True, **kwargs):
pass
@abstractmethod

View File

@@ -49,7 +49,7 @@ class ChunkCache(BasePrefixCache):
last_host_node=None,
)
def cache_finished_req(self, req: Req, insert: bool = True):
def cache_finished_req(self, req: Req, is_insert: bool = True):
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx,
# For decode server: if req.output_ids is empty, we want to free all req.origin_input_ids

View File

@@ -330,18 +330,18 @@ class RadixCache(BasePrefixCache):
return self._insert_helper(self.root_node, key, value)
def cache_finished_req(self, req: Req):
def cache_finished_req(self, req: Req, is_insert: bool = True):
"""Cache request when it finishes."""
all_token_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
if self.disable:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : len(req.origin_input_ids) + len(req.output_ids) - 1
req.req_pool_idx, :all_token_len
]
self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx)
return
token_ids = (req.origin_input_ids + req.output_ids)[:-1]
all_token_len = len(token_ids)
token_ids = (req.origin_input_ids + req.output_ids)[:all_token_len]
# For EAGLE radix cache, we will convert the key to bigram key, e.g. [1,2,3,4] -> [(1,2), (2,3), (3,4)], the length will -1. ((len([(1,2), (2,3), (3,4)]) = len([1,2,3,4]) - 1))
# So for the corresponding kv length should also -1. Then we get the actual_kv_len, and use it to do later calculation and slicing.
actual_kv_len = all_token_len - 1 if self.is_eagle else all_token_len
@@ -354,12 +354,9 @@ class RadixCache(BasePrefixCache):
page_aligned_kv_indices = kv_indices[:page_aligned_len].to(
dtype=torch.int64, copy=True
)
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
else:
page_aligned_len = actual_kv_len
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
if self.is_eagle:
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
page_aligned_token_len = (
page_aligned_len + 1 if self.is_eagle else page_aligned_len
@@ -372,11 +369,22 @@ class RadixCache(BasePrefixCache):
old_prefix_len -= 1
# Radix Cache takes one ref in memory pool
new_prefix_len = self.insert(
RadixKey(token_ids[:page_aligned_token_len], req.extra_key),
page_aligned_kv_indices,
)
self.token_to_kv_pool_allocator.free(kv_indices[old_prefix_len:new_prefix_len])
if is_insert:
new_prefix_len = self.insert(
RadixKey(token_ids[:page_aligned_token_len], req.extra_key),
page_aligned_kv_indices,
)
# Free the duplicates that were already in the tree
self.token_to_kv_pool_allocator.free(
kv_indices[old_prefix_len:new_prefix_len]
)
else:
self.token_to_kv_pool_allocator.free(
kv_indices[old_prefix_len:page_aligned_len]
)
# free the unaligned tail
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
# Remove req slot release the cache lock
self.req_to_token_pool.free(req.req_pool_idx)

View File

@@ -151,32 +151,37 @@ class RadixCacheCpp(BasePrefixCache):
def total_size(self):
return self.tree.total_size()
def cache_finished_req(self, req: Req):
def cache_finished_req(self, req: Req, is_insert: bool = True):
"""Cache request when it finishes."""
assert req.req_pool_idx is not None
token_ids = (req.origin_input_ids + req.output_ids)[:-1]
all_token_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
token_ids = (req.origin_input_ids + req.output_ids)[:all_token_len]
overall_len = len(token_ids) # prefill + decode
kv_indices = self.req_to_token_pool.req_to_token[req.req_pool_idx, :overall_len]
# NOTE: our C++ implementation don't need `token_ids` and `kv_indices` to be page-aligned
# it will automatically align them, but length of them should be equal
old_prefix_len = len(req.prefix_indices) // self.page_size * self.page_size
new_prefix_len = self._insert(RadixKey(token_ids, req.extra_key), kv_indices)
page_aligned_overall_len = overall_len // self.page_size * self.page_size
# NOTE: kv_indices[:old_prefix_len] == req.prefix_indices
assert old_prefix_len <= new_prefix_len, "Wrong prefix indices"
# KVCache between old & new is newly generated, but already exists in the pool
# we need to free this newly generated kv indices
if old_prefix_len < new_prefix_len:
self.token_to_kv_pool.free(kv_indices[old_prefix_len:new_prefix_len])
if is_insert:
new_prefix_len = self._insert(
RadixKey(token_ids, req.extra_key), kv_indices
)
# NOTE: kv_indices[:old_prefix_len] == req.prefix_indices
assert old_prefix_len <= new_prefix_len, "Wrong prefix indices"
# Free duplicates that were already in the pool
if old_prefix_len < new_prefix_len:
self.token_to_kv_pool.free(kv_indices[old_prefix_len:new_prefix_len])
else:
self.token_to_kv_pool.free(
kv_indices[old_prefix_len:page_aligned_overall_len]
)
# need to free the unaligned part, since it cannot be inserted into the radix tree
if self.page_size != 1 and ( # unaligned tail only exists when page_size > 1
(unaligned_len := overall_len % self.page_size) > 0
):
if page_aligned_overall_len < overall_len:
# NOTE: sglang PagedAllocator support unaligned free (which will automatically align it)
self.token_to_kv_pool.free(kv_indices[overall_len - unaligned_len :])
self.token_to_kv_pool.free(kv_indices[page_aligned_overall_len:])
# Remove req slot release the cache lock
self.dec_lock_ref(req.last_node)

View File

@@ -217,10 +217,12 @@ class LMCRadixCache(RadixCache):
return base_res
def cache_finished_req(self, req: "Req") -> None: # type: ignore[override]
def cache_finished_req(self, req: "Req", is_insert: bool = True) -> None: # type: ignore[override]
"""On request completion, insert device KV into radix and store to LMCache."""
super().cache_finished_req(req)
super().cache_finished_req(req, is_insert=is_insert)
if not is_insert:
return
token_ids = (req.origin_input_ids + req.output_ids)[:-1]
kv_indices = self.req_to_token_pool.req_to_token[

View File

@@ -427,19 +427,18 @@ class SWARadixCache(BasePrefixCache):
return self._insert_helper(self.root_node, key, value, prev_prefix_len)
def cache_finished_req(self, req: Req) -> None:
def cache_finished_req(self, req: Req, is_insert: bool = True) -> None:
"""Cache request when it finishes."""
all_token_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
if self.disable:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx,
: len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0),
req.req_pool_idx, :all_token_len
]
self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx)
return
token_ids = (req.origin_input_ids + req.output_ids)[:-1]
all_token_len = len(token_ids)
token_ids = (req.origin_input_ids + req.output_ids)[:all_token_len]
# For EAGLE radix cache, we will convert the key to bigram key, e.g. [1,2,3,4] -> [(1,2), (2,3), (3,4)], the length will -1. ((len([(1,2), (2,3), (3,4)]) = len([1,2,3,4]) - 1))
# So for the corresponding kv length should also -1. Then we get the actual_kv_len, and use it to do later calculation and slicing.
actual_kv_len = all_token_len - 1 if self.is_eagle else all_token_len
@@ -452,7 +451,6 @@ class SWARadixCache(BasePrefixCache):
page_aligned_kv_indices = kv_indices[:page_aligned_len].to(
dtype=torch.int64, copy=True
)
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
else:
page_aligned_len = actual_kv_len
page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True)
@@ -472,11 +470,19 @@ class SWARadixCache(BasePrefixCache):
# Radix Cache takes one ref in memory pool
# insert the token_ids and kv_indices into the radix tree
# Note: the insert function already frees the overlapped kv_indices
new_prefix_len = self.insert(
RadixKey(token_ids[:page_aligned_token_len], req.extra_key),
page_aligned_kv_indices,
old_prefix_len,
)
if is_insert:
new_prefix_len = self.insert(
RadixKey(token_ids[:page_aligned_token_len], req.extra_key),
page_aligned_kv_indices,
old_prefix_len,
)
else:
self.token_to_kv_pool_allocator.free(
kv_indices[old_prefix_len:page_aligned_len]
)
# free the unaligned tail
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
# Remove req slot release the cache lock
self.req_to_token_pool.free(req.req_pool_idx)