diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 7892ac940..a4c14d784 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -929,7 +929,11 @@ class SchedulerDisaggregationPrefillMixin: else: logger.warning(error_message) req.time_stats.trace_ctx.abort(abort_info={"reason": error_message}) - if req.req_pool_idx is not None or self.tree_cache.supports_mamba(): + if ( + req.req_pool_idx is not None + or req.kv is not None + or req.mamba_pool_idx is not None + ): release_kv_cache(req, self.tree_cache) maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator) req.pending_bootstrap = False diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index abad9d4aa..2667ba014 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2564,8 +2564,7 @@ class Scheduler( req.pending_bootstrap = False if self.enable_hicache_storage: self.tree_cache.release_aborted_request(req.rid) - if req.req_pool_idx is not None or self.tree_cache.supports_mamba(): - release_kv_cache(req, self.tree_cache, is_insert=False) + release_kv_cache(req, self.tree_cache, is_insert=False) self.chunked_req = None self._pending_chunked_abort_req = None diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 9722699bf..b9937b258 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -246,7 +246,7 @@ class SchedulerInvariantChecker: swa_uncached = 0 for batch in batches: for req in batch.reqs: - if req.req_pool_idx is None: + if req.kv is None: continue allocated_len = req.kv.kv_allocated_len diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 28cf9520b..a5d00d6a0 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -630,6 +630,8 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor: def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True): + # the two resources currently have the same lifecycle, thus simplify logic below + assert (req.req_pool_idx is None) == (req.kv is None) # MambaRadixCache may alloc mamba state before alloc KV cache if req.req_pool_idx is None: assert ( @@ -652,7 +654,8 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr # StreamingSession.cache_finished_req handles speculative tail trim # internally, then sets req_pool_idx = None. - if req.req_pool_idx is None: + assert (req.req_pool_idx is None) == (req.kv is None) + if req.req_pool_idx is None and req.kv is None: return start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len diff --git a/python/sglang/srt/session/streaming_session.py b/python/sglang/srt/session/streaming_session.py index 946aa438e..cb2a012e4 100644 --- a/python/sglang/srt/session/streaming_session.py +++ b/python/sglang/srt/session/streaming_session.py @@ -63,7 +63,7 @@ class SessionSlot: @property def is_holding_kv(self) -> bool: """Whether this slot currently holds KV pool resources.""" - return self.req_pool_idx is not None + return self.kv is not None def save_from_req(self, req: Req, is_first: bool): """Save KV state from a finishing request into this slot.""" @@ -212,7 +212,7 @@ class StreamingSession(BasePrefixCache): if not _is_streaming(req): return None slot = self.slots.get(req.session.session_id) - if slot is None or slot.req_pool_idx is None: + if slot is None or slot.kv is None: return None if req.to_finish is not None: req.session.abort_req() diff --git a/python/sglang/test/scripted_runtime/req_handle.py b/python/sglang/test/scripted_runtime/req_handle.py index a08c1aa8a..24f2cfd88 100644 --- a/python/sglang/test/scripted_runtime/req_handle.py +++ b/python/sglang/test/scripted_runtime/req_handle.py @@ -42,7 +42,7 @@ class ScriptedReqHandle: @property def kv_pages(self) -> int: req = self.req - if req is None or req.req_pool_idx is None: + if req is None or req.kv is None: return 0 page_size = self.context.scheduler.page_size return (req.kv.kv_allocated_len + page_size - 1) // page_size