Fix abusing presence of req.req_pool_idx to indicate the presence of req.kv resources (#29430)
This commit is contained in:
@@ -929,7 +929,11 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
else:
|
else:
|
||||||
logger.warning(error_message)
|
logger.warning(error_message)
|
||||||
req.time_stats.trace_ctx.abort(abort_info={"reason": 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)
|
release_kv_cache(req, self.tree_cache)
|
||||||
maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator)
|
maybe_release_metadata_buffer(req, self.req_to_metadata_buffer_idx_allocator)
|
||||||
req.pending_bootstrap = False
|
req.pending_bootstrap = False
|
||||||
|
|||||||
@@ -2564,8 +2564,7 @@ class Scheduler(
|
|||||||
req.pending_bootstrap = False
|
req.pending_bootstrap = False
|
||||||
if self.enable_hicache_storage:
|
if self.enable_hicache_storage:
|
||||||
self.tree_cache.release_aborted_request(req.rid)
|
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.chunked_req = None
|
||||||
self._pending_chunked_abort_req = None
|
self._pending_chunked_abort_req = None
|
||||||
|
|||||||
@@ -246,7 +246,7 @@ class SchedulerInvariantChecker:
|
|||||||
swa_uncached = 0
|
swa_uncached = 0
|
||||||
for batch in batches:
|
for batch in batches:
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
if req.req_pool_idx is None:
|
if req.kv is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
allocated_len = req.kv.kv_allocated_len
|
allocated_len = req.kv.kv_allocated_len
|
||||||
|
|||||||
@@ -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):
|
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
|
# MambaRadixCache may alloc mamba state before alloc KV cache
|
||||||
if req.req_pool_idx is None:
|
if req.req_pool_idx is None:
|
||||||
assert (
|
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
|
# StreamingSession.cache_finished_req handles speculative tail trim
|
||||||
# internally, then sets req_pool_idx = None.
|
# 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
|
return
|
||||||
|
|
||||||
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ class SessionSlot:
|
|||||||
@property
|
@property
|
||||||
def is_holding_kv(self) -> bool:
|
def is_holding_kv(self) -> bool:
|
||||||
"""Whether this slot currently holds KV pool resources."""
|
"""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):
|
def save_from_req(self, req: Req, is_first: bool):
|
||||||
"""Save KV state from a finishing request into this slot."""
|
"""Save KV state from a finishing request into this slot."""
|
||||||
@@ -212,7 +212,7 @@ class StreamingSession(BasePrefixCache):
|
|||||||
if not _is_streaming(req):
|
if not _is_streaming(req):
|
||||||
return None
|
return None
|
||||||
slot = self.slots.get(req.session.session_id)
|
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
|
return None
|
||||||
if req.to_finish is not None:
|
if req.to_finish is not None:
|
||||||
req.session.abort_req()
|
req.session.abort_req()
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ class ScriptedReqHandle:
|
|||||||
@property
|
@property
|
||||||
def kv_pages(self) -> int:
|
def kv_pages(self) -> int:
|
||||||
req = self.req
|
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
|
return 0
|
||||||
page_size = self.context.scheduler.page_size
|
page_size = self.context.scheduler.page_size
|
||||||
return (req.kv.kv_allocated_len + page_size - 1) // page_size
|
return (req.kv.kv_allocated_len + page_size - 1) // page_size
|
||||||
|
|||||||
Reference in New Issue
Block a user