[Scheduler] Count the parked chunked-prefill request in the busy mem check (#37502)

This commit is contained in:
Zhiqiang Xie
2026-09-03 12:08:17 -07:00
committed by GitHub
parent 2da5802bfa
commit 68978f8d52
2 changed files with 26 additions and 15 deletions
+1
View File
@@ -2293,6 +2293,7 @@ class Scheduler(
pool_stats_observer=self.pool_stats_observer,
get_last_batch=lambda: self.last_batch,
get_running_batch=lambda: self.running_batch,
get_chunked_req=lambda: self.chunked_req,
)
def init_rank_consensus_checker(self) -> None:
@@ -58,6 +58,9 @@ class SchedulerInvariantChecker:
pool_stats_observer: SchedulerPoolStatsObserver
get_last_batch: Callable
get_running_batch: Callable
# The chunked-prefill request parked between chunks is in neither batch;
# its uncached tokens must still be counted.
get_chunked_req: Callable = field(default=lambda: None)
count_req_pool_leak_warnings: int = 0
count_memory_leak_warnings: int = 0
recent_busy_msgs: Deque[str] = field(
@@ -262,24 +265,31 @@ class SchedulerInvariantChecker:
full_uncached = 0
swa_uncached = 0
for batch in batches:
for req in batch.reqs:
if not req.kv.holds_kv:
continue
counted: set[int] = set()
reqs = [req for batch in batches for req in batch.reqs]
chunked_req = self.get_chunked_req()
if chunked_req is not None:
reqs.append(chunked_req)
for req in reqs:
if id(req) in counted:
continue
counted.add(id(req))
if not req.kv.holds_kv:
continue
allocated_len = req.kv.kv_allocated_len
if self.page_size > 1:
allocated_len = ceil_align(allocated_len, self.page_size)
assert req.kv.cache_protected_len % self.page_size == 0
allocated_len = req.kv.kv_allocated_len
if self.page_size > 1:
allocated_len = ceil_align(allocated_len, self.page_size)
assert req.kv.cache_protected_len % self.page_size == 0
full_uncached += allocated_len - req.kv.cache_protected_len
if self.is_hybrid_swa:
swa_uncached += allocated_len - max(
req.kv.cache_protected_len, req.kv.swa_evicted_seqlen
)
full_uncached += allocated_len - req.kv.cache_protected_len
if self.is_hybrid_swa:
swa_uncached += allocated_len - max(
req.kv.cache_protected_len, req.kv.swa_evicted_seqlen
)
if req.beam_group is not None:
full_uncached += req.beam_group.extra_uncached_tokens()
if req.beam_group is not None:
full_uncached += req.beam_group.extra_uncached_tokens()
return full_uncached, swa_uncached