diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 51167bc60..9d9e09e66 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1445,7 +1445,7 @@ class Req(ReqDllmMixin): kv, self.kv = self.kv, ReqKvInfo() return kv - def effective_kv_committed_len(self) -> int: + def owned_kv_len(self) -> int: # Report only the prompt prefix so thinking + answer fall into the # overallocated range and are reclaimed by release_kv_cache. #22373. if get_serving().strip_thinking_cache and self.reasoning_tokens > 0: diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index 5d260f241..81a1a5d95 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -434,8 +434,18 @@ class BasePrefixCache(ABC, PrefixCacheTrait): return None @abstractmethod - def cache_finished_req(self, req: Req, is_insert: bool = True, **kwargs): - pass + def cache_finished_req( + self, req: Req, is_insert: bool = True, *, owned_kv_len: int, **kwargs + ): + """Dispose of a finished request's KV. + + ``[0, req.kv.cache_protected_len)`` is cache-owned and must survive. + Every slot in ``[req.kv.cache_protected_len, owned_kv_len)`` is this + call's to account for: insert what can be keyed, release the rest. + Slicing the kv row by the token-id count instead strands whatever + lies between -- no caller releases those. ``release_kv_cache`` frees + everything past ``owned_kv_len``. + """ @abstractmethod def cache_unfinished_req(self, req: Req, **kwargs): diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index 75c54a074..fbfca879a 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -30,7 +30,6 @@ class CacheInitParams: # Keyword arguments for the eviction policy's constructor; see the strategy # classes in evict_policy.py for what each policy accepts. eviction_policy_config: Optional[dict[str, Any]] = None - disable_finished_insert: bool = False enable_metrics: bool = False enable_kv_cache_events: bool = False diff --git a/python/sglang/srt/mem_cache/chunk_cache.py b/python/sglang/srt/mem_cache/chunk_cache.py index ab3bc4afb..8c005ac39 100644 --- a/python/sglang/srt/mem_cache/chunk_cache.py +++ b/python/sglang/srt/mem_cache/chunk_cache.py @@ -77,11 +77,11 @@ class ChunkCache(BasePrefixCache): return InsertResult(prefix_len=0) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ): # For decode server: if req.output_ids is empty, we want to free all req.origin_input_ids # The protected prefix is not this req's to free. - self.free_kv_row(req.kv, [(req.kv.cache_protected_len, kv_len_to_handle)]) + self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)]) def cache_unfinished_req(self, req: Req, chunked=False): kv_indices = self.req_to_token_pool.req_to_token[ @@ -151,11 +151,10 @@ class PureSWAChunkCache(SWAChunkCache): """ def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ): - kv_committed_len = kv_len_to_handle kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_committed_len + req.kv.req_pool_idx, :owned_kv_len ] # The cache_protected_len prefix is not this req's to free. protected_len = req.kv.cache_protected_len @@ -165,9 +164,9 @@ class PureSWAChunkCache(SWAChunkCache): parts = [] if evict_floor > protected_len: parts.append(kv_indices[protected_len:evict_floor]) - if evicted_seqlen < kv_committed_len: + if evicted_seqlen < owned_kv_len: parts.append( - kv_indices[max(evicted_seqlen, protected_len) : kv_committed_len] + kv_indices[max(evicted_seqlen, protected_len) : owned_kv_len] ) if parts: self.token_to_kv_pool_allocator.free(torch.cat(parts)) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 22dd0729a..fda235acf 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -284,11 +284,11 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr req.kv.mamba_pool_idx = None return - effective_kv_committed_len = req.effective_kv_committed_len() + owned_kv_len = req.owned_kv_len() tree_cache.cache_finished_req( req, is_insert=is_insert and not getattr(req, "skip_radix_cache_insert", False), - kv_len_to_handle=effective_kv_committed_len, + owned_kv_len=owned_kv_len, ) # StreamingSession.cache_finished_req handles speculative tail trim @@ -297,7 +297,7 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr if not req.kv.holds_kv: return - start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len + start_p, end_p = owned_kv_len, req.kv.kv_allocated_len _release_overallocated_kv_indices(req, start_p, end_p, tree_cache) # If the prefix cache doesn't manage mamba states, we must free them here. diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index e98bdec52..e2faa6efe 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -544,20 +544,20 @@ class MambaRadixCache(BasePrefixCache): return InsertResult(prefix_len=prefix_len, mamba_exist=mamba_exist) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ) -> None: """Cache request when it finishes.""" if self.disable: kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_len_to_handle + req.kv.req_pool_idx, :owned_kv_len ] self.token_to_kv_pool_allocator.free_segment(kv_indices, start_pos=0) self.req_to_token_pool.free_mamba_cache(req) return - token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_len_to_handle + req.kv.req_pool_idx, :owned_kv_len ] if is_insert: @@ -607,7 +607,7 @@ class MambaRadixCache(BasePrefixCache): page_aligned_kv_indices = kv_indices.to(dtype=torch.int64, copy=True) assert cache_len == page_aligned_len, ( - f"It is required {cache_len=}, {page_aligned_len=}, {kv_len_to_handle=}, {len(req.origin_input_ids)=}, {len(req.output_ids)=} ping @yizhang2077 if you see this" + f"It is required {cache_len=}, {page_aligned_len=}, {owned_kv_len=}, {len(req.origin_input_ids)=}, {len(req.output_ids)=} ping @yizhang2077 if you see this" ) # Radix Cache takes one ref in memory pool diff --git a/python/sglang/srt/mem_cache/pure_swa_radix_cache.py b/python/sglang/srt/mem_cache/pure_swa_radix_cache.py index 0636d0c13..2a652fc31 100644 --- a/python/sglang/srt/mem_cache/pure_swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/pure_swa_radix_cache.py @@ -63,7 +63,7 @@ class PureSWARadixCache(RadixCache): return super().evict(EvictParams(num_tokens=num_tokens)) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ): """Cache request when it finishes. @@ -72,20 +72,16 @@ class PureSWARadixCache(RadixCache): to the allocator. The range [evict_floor, swa_evicted_seqlen) was already freed by _evict_swa during decode — we skip it to avoid double-free. """ - if self.disable_finished_insert: - is_insert = False - - kv_committed_len = kv_len_to_handle if self.disable: kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_committed_len + req.kv.req_pool_idx, :owned_kv_len ] self.token_to_kv_pool_allocator.free(kv_indices) return - token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_committed_len + req.kv.req_pool_idx, :owned_kv_len ] radix_key = RadixKey( diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index 05c31b6c3..3eaebc5e9 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -325,7 +325,6 @@ class RadixCache(BasePrefixCache): self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator self.page_size = params.page_size self.is_eagle = params.is_eagle - self.disable_finished_insert = params.disable_finished_insert self.eviction_policy = params.eviction_policy.lower() self.kv_events = KVCacheEventRecorder( @@ -477,17 +476,13 @@ class RadixCache(BasePrefixCache): return InsertResult(prefix_len=prefix_len, last_device_node=last_node) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ): """Cache request when it finishes.""" - # In deterministic mode, disable finished request insertion to radix cache - if self.disable_finished_insert: - is_insert = False - if self.disable: # The protected prefix is not this req's to free. kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, req.kv.cache_protected_len : kv_len_to_handle + req.kv.req_pool_idx, req.kv.cache_protected_len : owned_kv_len ] self.token_to_kv_pool_allocator.free_segment( kv_indices, start_pos=req.kv.cache_protected_len @@ -498,7 +493,7 @@ class RadixCache(BasePrefixCache): # Frees committed slots that no token id names, which the insert # path below cannot reach; the protected prefix stays with the cache. kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, req.kv.cache_protected_len : kv_len_to_handle + req.kv.req_pool_idx, req.kv.cache_protected_len : owned_kv_len ] self.token_to_kv_pool_allocator.free_segment( kv_indices, start_pos=req.kv.cache_protected_len @@ -507,9 +502,9 @@ class RadixCache(BasePrefixCache): self.dec_lock_ref(req.last_node) return - token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, : len(token_ids) + req.kv.req_pool_idx, :owned_kv_len ] radix_key = RadixKey( diff --git a/python/sglang/srt/mem_cache/radix_cache_cpp.py b/python/sglang/srt/mem_cache/radix_cache_cpp.py index d1dc22ba5..168df6eb1 100644 --- a/python/sglang/srt/mem_cache/radix_cache_cpp.py +++ b/python/sglang/srt/mem_cache/radix_cache_cpp.py @@ -181,20 +181,20 @@ class RadixCacheCpp(BasePrefixCache): return self.tree.total_size() def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ): """Cache request when it finishes.""" self._reject_cache_salt(req.cache_salt) assert req.kv.holds_kv - token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_len_to_handle + req.kv.req_pool_idx, :owned_kv_len ].to(dtype=torch.int64, copy=True) # 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 - page_aligned_overall_len = kv_len_to_handle // self.page_size * self.page_size + page_aligned_overall_len = owned_kv_len // self.page_size * self.page_size if is_insert: new_prefix_len = self._insert( @@ -213,7 +213,7 @@ class RadixCacheCpp(BasePrefixCache): ) # need to free the unaligned part, since it cannot be inserted into the radix tree - if page_aligned_overall_len < kv_len_to_handle: + if page_aligned_overall_len < owned_kv_len: # NOTE: sglang PagedAllocator support unaligned free (which will automatically align it) self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_overall_len:]) diff --git a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py index deba6a699..0a14cd84a 100644 --- a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py @@ -386,12 +386,10 @@ class FlexKVRadixCache(RadixCache): # ------------------------------------------------------------------ def cache_finished_req( # type: ignore[override] - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ) -> None: """Base cache_finished_req then fire an async FlexKV store.""" - super().cache_finished_req( - req, is_insert=is_insert, kv_len_to_handle=kv_len_to_handle - ) + super().cache_finished_req(req, is_insert=is_insert, owned_kv_len=owned_kv_len) if not is_insert: self._load_markers.pop(req.cache_request_handle, None) return diff --git a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py index e07754690..fdb019bb4 100644 --- a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py @@ -439,13 +439,11 @@ class LMCRadixCache(RadixCache): ) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ) -> None: """On request completion, insert device KV into radix and store to LMCache.""" - super().cache_finished_req( - req, is_insert=is_insert, kv_len_to_handle=kv_len_to_handle - ) + super().cache_finished_req(req, is_insert=is_insert, owned_kv_len=owned_kv_len) if not is_insert: if self._mode is LMCacheMode.MP: self._mp_load_back_markers.pop(req.rid, None) diff --git a/python/sglang/srt/mem_cache/swa_radix_cache.py b/python/sglang/srt/mem_cache/swa_radix_cache.py index dbc6d8b91..ffce936f0 100644 --- a/python/sglang/srt/mem_cache/swa_radix_cache.py +++ b/python/sglang/srt/mem_cache/swa_radix_cache.py @@ -460,16 +460,16 @@ class SWARadixCache(BasePrefixCache): return InsertResult(prefix_len=prefix_len) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int + self, req: Req, is_insert: bool = True, *, owned_kv_len: int ) -> None: """Cache request when it finishes.""" if self.disable: - self.free_kv_row(req.kv, [(0, kv_len_to_handle)]) + self.free_kv_row(req.kv, [(0, owned_kv_len)]) return - token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_len_to_handle + req.kv.req_pool_idx, :owned_kv_len ] radix_key = RadixKey( @@ -497,7 +497,7 @@ class SWARadixCache(BasePrefixCache): self.free_kv_row(req.kv, [(old_prefix_len, page_aligned_len)]) # free the unaligned tail - self.free_kv_row(req.kv, [(page_aligned_len, kv_len_to_handle)]) + self.free_kv_row(req.kv, [(page_aligned_len, owned_kv_len)]) # Remove req slot release the cache lock self.dec_lock_ref( diff --git a/python/sglang/srt/mem_cache/unified_cache/components/README.md b/python/sglang/srt/mem_cache/unified_cache/components/README.md index f0b1c61b6..c5d0e5cf1 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/README.md +++ b/python/sglang/srt/mem_cache/unified_cache/components/README.md @@ -228,14 +228,14 @@ receipt proves were taken. The eventual full release must pass --- -### `cache_finished_req(req: Req, is_insert: bool = True, *, kv_len_to_handle: int)` +### `cache_finished_req(req: Req, is_insert: bool = True, *, owned_kv_len: int)` Cache a completed request's KV data into the tree. | Aspect | Detail | |--------|--------| | **Purpose** | After a request finishes, insert its token/KV data into the tree for future reuse | -| **Inputs** | `req` — the finished request; `is_insert` — whether to insert (True) or just release locks (False); `kv_len_to_handle` — committed KV length supplied by the caller | +| **Inputs** | `req` — the finished request; `is_insert` — insert the owned range into the tree (True) or free it (False); `owned_kv_len` — end of the request-owned KV range; slots past it are freed by `release_kv_cache` | | **Output** | `None` | | **Mutation** | Calls component hooks → `insert` → `dec_lock_ref` → component cleanup. Frees unaligned tail KV indices; frees non-inserted KV indices when `is_insert=False`. | | **Complexity** | **O(K + D·C)** — insert O(K + D·C) + lock release O(D). Simplifies to **O(K)**. | diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 7b08faa6c..ad6d463ca 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -955,22 +955,22 @@ class UnifiedRadixCache(BasePrefixCache): return DecLockRefResult() return self.tree_core.dec_host_lock_ref(node_id, params) - @rank_consensus(same_params=["req.rid", "is_insert", "kv_len_to_handle"]) + @rank_consensus(same_params=["req.rid", "is_insert", "owned_kv_len"]) def cache_finished_req( - self, req: Req, is_insert: bool = True, *, kv_len_to_handle: int, **kwargs + self, req: Req, is_insert: bool = True, *, owned_kv_len: int, **kwargs ) -> None: if self.session.try_cache_finished_req(req, is_insert=is_insert, **kwargs): return if self.disable: - self.free_kv_row(req.kv, [(0, kv_len_to_handle)]) + self.free_kv_row(req.kv, [(0, owned_kv_len)]) for comp in self._components_tuple: comp.cleanup_after_caching_req(req, is_finished=True) return - token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] + token_ids = (req.origin_input_ids + req.output_ids)[:owned_kv_len] kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, :kv_len_to_handle + req.kv.req_pool_idx, :owned_kv_len ] result = None @@ -1078,7 +1078,7 @@ class UnifiedRadixCache(BasePrefixCache): ranges.append((tail_free_start, len(kv_indices_full))) self.free_kv_row(req.kv, ranges) else: - self.free_kv_row(req.kv, [(req.kv.cache_protected_len, kv_len_to_handle)]) + self.free_kv_row(req.kv, [(req.kv.cache_protected_len, owned_kv_len)]) # Synthetic profiling requests may own KV without locking a tree node. if req.last_node is not None: diff --git a/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py b/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py index c366648e4..89d094418 100644 --- a/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py +++ b/test/registered/unit/disaggregation/test_specv2_kvcache_offloading.py @@ -55,7 +55,7 @@ def _make_mock_req( kv_allocated_len=kv_allocated_len, ) req.prefix_indices = list(range(prefix_indices_len)) - req.effective_kv_committed_len = lambda: req.kv.kv_committed_len + req.owned_kv_len = lambda: req.kv.kv_committed_len return req diff --git a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py index af68e287b..4192dcf3a 100644 --- a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py +++ b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py @@ -234,7 +234,7 @@ class TestDecodeLockRefScenarios(CustomTestCase): cache.cache_unfinished_req(req) # Step 3: cache_finished_req with is_insert=True (dec lock) - cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len) + cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len) # Verify: all non-root nodes should have lock_ref == 0 # (root always has lock_ref == 1) @@ -283,7 +283,7 @@ class TestDecodeLockRefScenarios(CustomTestCase): cache.cache_unfinished_req(req) # Step 3: cache_finished_req (dec leaf) - cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len) + cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len) # Root lock unchanged, all nodes unlocked self.assertEqual(cache.root_node.lock_ref, root_lock_before) @@ -331,7 +331,7 @@ class TestDecodeLockRefScenarios(CustomTestCase): # Transfer fails -> cache_finished_req with is_insert=False cache.token_to_kv_pool_allocator.reset_mock() cache.cache_finished_req( - req, is_insert=False, kv_len_to_handle=req.kv.kv_committed_len + req, is_insert=False, owned_kv_len=req.kv.kv_committed_len ) free_call = cache.token_to_kv_pool_allocator.free_segment.call_args @@ -346,6 +346,53 @@ class TestDecodeLockRefScenarios(CustomTestCase): # Prefix tokens should still be in tree and evictable self.assertEqual(cache.evictable_size(), len(prefix)) + def test_insert_releases_committed_slot_without_token_id(self): + """The insert path releases up to owned_kv_len, not len(token_ids). + + Pins the ownership contract on BasePrefixCache.cache_finished_req. + """ + cache, req_to_token = _make_cache_with_pools() + + prefix = [1, 2, 3] + prefix_vals = [10, 20, 30] + self._populate_prefix(cache, prefix, prefix_vals) + + result = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", prefix)))) + matched_node = result.last_device_node + prefix_len = len(result.device_indices) + cache.inc_lock_ref(matched_node) + + # Token sequence is 5 long; a 6th KV slot is committed with no token id. + full_ids = [1, 2, 3, 4, 5] + row_vals = [10, 20, 30, 40, 50, 60] + req_to_token[0, : len(row_vals)] = torch.tensor(row_vals, dtype=torch.int64) + + req = _make_req( + fill_ids=full_ids, + req_pool_idx=0, + cache_protected_len=prefix_len, + last_node=matched_node, + ) + req.kv.kv_committed_len = len(row_vals) + req.kv.kv_allocated_len = len(row_vals) + + cache.token_to_kv_pool_allocator.reset_mock() + cache.cache_finished_req( + req, is_insert=True, owned_kv_len=req.kv.kv_committed_len + ) + + # The unnamed tail slot is freed as the segment past the radix key. + segments = cache.token_to_kv_pool_allocator.free_segments.call_args.args[0] + freed = { + int(start_pos) + i: int(v) + for indices, start_pos in segments + for i, v in enumerate(indices.tolist()) + } + self.assertIn( + len(full_ids), freed, "committed slot with no token id was not freed" + ) + self.assertEqual(freed[len(full_ids)], row_vals[-1]) + def test_full_transfer_failure(self): """Scenario 4: no prefix match, transfer fails. @@ -384,7 +431,7 @@ class TestDecodeLockRefScenarios(CustomTestCase): # Transfer fails -> cache_finished_req with is_insert=False # dec_lock_ref(root) is a no-op cache.cache_finished_req( - req, is_insert=False, kv_len_to_handle=req.kv.kv_committed_len + req, is_insert=False, owned_kv_len=req.kv.kv_committed_len ) # Root lock unchanged, nothing protected or evictable @@ -572,7 +619,7 @@ class TestDecodeLockRefScenarios(CustomTestCase): ) cache.cache_unfinished_req(req) - cache.cache_finished_req(req, kv_len_to_handle=req.kv.kv_committed_len) + cache.cache_finished_req(req, owned_kv_len=req.kv.kv_committed_len) # After all iterations, root lock should be 1, no protected nodes self.assertEqual(cache.root_node.lock_ref, 1) diff --git a/test/registered/unit/mem_cache/test_page_interleave_shard.py b/test/registered/unit/mem_cache/test_page_interleave_shard.py index 508204574..6961f75fc 100644 --- a/test/registered/unit/mem_cache/test_page_interleave_shard.py +++ b/test/registered/unit/mem_cache/test_page_interleave_shard.py @@ -768,7 +768,7 @@ class TestRotationGraftDecline(CustomTestCase): req = _GraftReq(list(range(8)) + [90, 91, 92, 93]) req.kv_rotation_base = 3 own_locs = self._own_row(tree, req, 12) - tree.cache_finished_req(req, kv_len_to_handle=12) + tree.cache_finished_req(req, owned_kv_len=12) released = torch.cat(freed) # Everything past the protected prefix is released: the duplicates of # the matched region AND the declined tail (nothing leaks, nothing is @@ -784,7 +784,7 @@ class TestRotationGraftDecline(CustomTestCase): req = _GraftReq(list(range(8)) + [90, 91, 92, 93]) req.kv_rotation_base = 1 own_locs = self._own_row(tree, req, 12) - tree.cache_finished_req(req, kv_len_to_handle=12) + tree.cache_finished_req(req, owned_kv_len=12) self.assertEqual(_match_len(tree, req.fill_ids), 12) released = torch.cat(freed) if freed else torch.empty(0, dtype=torch.int64) # Only the 8 duplicate rows go back; the tail stays live in the tree. diff --git a/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py b/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py index 39acb211c..6fd818a16 100644 --- a/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py +++ b/test/registered/unit/mem_cache/test_pure_swa_chunk_cache.py @@ -45,7 +45,7 @@ class TestPureSWAChunkCache(CustomTestCase): def test_finished_req_skips_already_evicted_swa_range(self): cache = self._make_cache() - cache.cache_finished_req(_FakeReq(), kv_len_to_handle=8) + cache.cache_finished_req(_FakeReq(), owned_kv_len=8) self.assertEqual(len(cache.token_to_kv_pool_allocator.freed), 1) freed = cache.token_to_kv_pool_allocator.freed[0] @@ -56,7 +56,7 @@ class TestPureSWAChunkCache(CustomTestCase): req = _FakeReq() req.kv.cache_protected_len = 2 - cache.cache_finished_req(req, kv_len_to_handle=8) + cache.cache_finished_req(req, owned_kv_len=8) freed = cache.token_to_kv_pool_allocator.freed[0] self.assertTrue(torch.equal(freed, torch.tensor([2, 6, 7]))) diff --git a/test/registered/unit/mem_cache/test_pure_swa_radix_cache.py b/test/registered/unit/mem_cache/test_pure_swa_radix_cache.py index b9f108dd1..25d3911cf 100644 --- a/test/registered/unit/mem_cache/test_pure_swa_radix_cache.py +++ b/test/registered/unit/mem_cache/test_pure_swa_radix_cache.py @@ -26,7 +26,6 @@ class TestPureSWARadixCache(CustomTestCase): def test_no_insert_frees_window_after_evict_floor_before_swa_eviction(self): allocator = _FakeAllocator() cache = PureSWARadixCache.__new__(PureSWARadixCache) - cache.disable_finished_insert = False cache.disable = False cache.is_eagle = False cache.page_size = 1 @@ -48,7 +47,7 @@ class TestPureSWARadixCache(CustomTestCase): ), ) - cache.cache_finished_req(req, is_insert=False, kv_len_to_handle=8) + cache.cache_finished_req(req, is_insert=False, owned_kv_len=8) self.assertEqual(allocator.freed, [[0, 1, 2, 3], [4, 5, 6, 7]]) diff --git a/test/registered/unit/mem_cache/test_radix_cache_unit.py b/test/registered/unit/mem_cache/test_radix_cache_unit.py index 29e94d89d..49319dfcf 100644 --- a/test/registered/unit/mem_cache/test_radix_cache_unit.py +++ b/test/registered/unit/mem_cache/test_radix_cache_unit.py @@ -572,7 +572,7 @@ class TestRadixCache(CustomTestCase): cache.cache_finished_req( req, is_insert=True, - kv_len_to_handle=len(prompt_ids) + len(output_ids), + owned_kv_len=len(prompt_ids) + len(output_ids), ) (prompt_node,) = cache.root_node.children.values() diff --git a/test/registered/unit/mem_cache/test_swa_eviction_boundary.py b/test/registered/unit/mem_cache/test_swa_eviction_boundary.py index 7462c2cf7..e486db1fd 100644 --- a/test/registered/unit/mem_cache/test_swa_eviction_boundary.py +++ b/test/registered/unit/mem_cache/test_swa_eviction_boundary.py @@ -191,7 +191,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): self.assertLess(req.kv.swa_evicted_seqlen, insert_len) tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len + req, is_insert=True, owned_kv_len=req._kv_committed_len ) tree.sanity_check() @@ -358,7 +358,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): ) tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len + req, is_insert=True, owned_kv_len=req._kv_committed_len ) tree.sanity_check() @@ -397,7 +397,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): pool.write((0, slice(0, first_len)), kv1) req1 = _make_req(0, list(range(first_len)), 0, tree) tree.cache_finished_req( - req1, is_insert=True, kv_len_to_handle=req1._kv_committed_len + req1, is_insert=True, owned_kv_len=req1._kv_committed_len ) tree.sanity_check() @@ -415,7 +415,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): swa_evictable_before = tree.swa_evictable_size_ tree.cache_finished_req( - req2, is_insert=True, kv_len_to_handle=req2._kv_committed_len + req2, is_insert=True, owned_kv_len=req2._kv_committed_len ) # New tokens [16, 24) should all be non-tombstone @@ -447,9 +447,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): self.assertGreater(req.kv.swa_evicted_seqlen, 0, "Should have some eviction") self.assertLess(req.kv.swa_evicted_seqlen, insert_len, "Should be partial") - tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len - ) + tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len) non_tombstone = insert_len - req.kv.swa_evicted_seqlen self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + non_tombstone) @@ -487,9 +485,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): req.kv.swa_evicted_seqlen = old_evicted swa_evictable_before = tree.swa_evictable_size_ - tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len - ) + tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len) self.assertEqual(tree.swa_evictable_size_, swa_evictable_before) @@ -519,7 +515,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): self.assertLess(req.kv.swa_evicted_seqlen, insert_len, f"turn {turn}") tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len + req, is_insert=True, owned_kv_len=req._kv_committed_len ) tree.sanity_check() @@ -545,7 +541,7 @@ class TestSWAEvictionBoundary(unittest.TestCase): ) tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len + req, is_insert=True, owned_kv_len=req._kv_committed_len ) tree.sanity_check() diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 65c9b815e..6b0cca4a8 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -813,9 +813,7 @@ class TestSWA(unittest.TestCase): return original_insert(params) tree.insert = wrapped_insert - tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req._kv_committed_len - ) + tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len) self.assertEqual(captured["prev_prefix_len"], req.kv.cache_protected_len) self.assertTrue(captured["is_bigram"]) @@ -849,7 +847,7 @@ class TestSWA(unittest.TestCase): allocator.free_segment = wrapped_free_segment tree.cache_finished_req( - req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len + req2, is_insert=False, owned_kv_len=req2._kv_committed_len ) # EAGLE + page_size=1 => page_aligned_len = committed_len - 1 = 5 @@ -1465,7 +1463,7 @@ class TestCacheUnfinishedReqEvictedPrefix(CustomTestCase): # Finishing drops the locks, which sanity_check needs; the accounting # must survive the re-walk. - tree.cache_finished_req(req, kv_len_to_handle=num_tokens) + tree.cache_finished_req(req, owned_kv_len=num_tokens) self.assertEqual(allocator.swa_available_size(), swa_before) self.assertEqual( tree.swa_evictable_size_ + tree.swa_protected_size_, diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py index 95460dc3c..c88f6e9a2 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py @@ -662,7 +662,7 @@ def bench_cache_finished( "cache_finished", lambda: req_items, lambda req: env.tree.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req.kv.kv_committed_len + req, is_insert=True, owned_kv_len=req.kv.kv_committed_len ), len(req_items) - warmup, env.avg_tokens, diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 08b3a097c..749b547a7 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -1668,9 +1668,7 @@ class UnifiedRadixCacheSuite: if self.cfg.has_mamba: req.kv.mamba_last_track_seqlen = kv_len - cache.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len() - ) + cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len()) all_ids = input_ids + output_ids aligned_len = (len(all_ids) // ps) * ps @@ -1732,9 +1730,9 @@ class UnifiedRadixCacheSuite: with get_serving().override(strip_thinking_cache=True): avail_before = allocator.available_size() cache.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len() + req, is_insert=True, owned_kv_len=req.owned_kv_len() ) - start_p, end_p = req.effective_kv_committed_len(), req.kv.kv_allocated_len + start_p, end_p = req.owned_kv_len(), req.kv.kv_allocated_len if ps > 1: start_p = ((start_p + ps - 1) // ps) * ps if start_p < end_p: @@ -1775,9 +1773,7 @@ class UnifiedRadixCacheSuite: ) avail_before = allocator.available_size() - cache.cache_finished_req( - req, is_insert=False, kv_len_to_handle=req.effective_kv_committed_len() - ) + cache.cache_finished_req(req, is_insert=False, owned_kv_len=req.owned_kv_len()) self.assertEqual(allocator.available_size(), avail_before + kv_len) m = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens)))) @@ -1952,9 +1948,7 @@ class UnifiedRadixCacheSuite: req.kv.mamba_last_track_seqlen = kv_len avail_before = allocator.available_size() - cache.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len() - ) + cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len()) self.assertEqual(allocator.available_size(), avail_before + tail_extra) aligned = input_ids[: (len(input_ids) // ps) * ps] @@ -8408,9 +8402,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase): ) req.last_node = cache.root_node_handle() - cache.cache_finished_req( - req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len() - ) + cache.cache_finished_req(req, is_insert=True, owned_kv_len=req.owned_kv_len()) def test_finished_req_stores_radix_mamba_state_in_int8_pool(self): cache, allocator, req_to_token_pool = build_fixture(self.cfg) diff --git a/test/registered/unit/mem_cache/test_unified_radix_lock_ref.py b/test/registered/unit/mem_cache/test_unified_radix_lock_ref.py index 8899d2820..0feb87d6a 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_lock_ref.py +++ b/test/registered/unit/mem_cache/test_unified_radix_lock_ref.py @@ -35,7 +35,7 @@ class TestUnifiedRadixLockRefScenarios(unittest.TestCase): swa_prefix_lock_released=False, ) - cache.cache_finished_req(req, is_insert=False, kv_len_to_handle=3) + cache.cache_finished_req(req, is_insert=False, owned_kv_len=3) cache.free_kv_row.assert_called_once_with(kv, [(0, 3)]) cache._dec_req_lock.assert_not_called() diff --git a/test/registered/xpu/test_lmcache_radix_cache.py b/test/registered/xpu/test_lmcache_radix_cache.py index fb2ee18f9..bf8842b19 100644 --- a/test/registered/xpu/test_lmcache_radix_cache.py +++ b/test/registered/xpu/test_lmcache_radix_cache.py @@ -193,7 +193,7 @@ class TestLMCRadixCacheXPU(unittest.TestCase): gt_v.append(v.clone()) req = _make_req("req-0", req_pool_idx, token_ids, tree) - tree.cache_finished_req(req, kv_len_to_handle=len(token_ids)) + tree.cache_finished_req(req, owned_kv_len=len(token_ids)) # IP-mode store is async on tree.store_stream; evict()'s # synchronize() is what the real scheduler relies on to make the # store visible before slots are reused.