From 68978f8d52547f83af76b62e5020e4dbfe0edf13 Mon Sep 17 00:00:00 2001 From: Zhiqiang Xie Date: Thu, 3 Sep 2026 12:08:17 -0700 Subject: [PATCH] [Scheduler] Count the parked chunked-prefill request in the busy mem check (#37502) --- python/sglang/srt/managers/scheduler.py | 1 + .../scheduler_components/invariant_checker.py | 40 ++++++++++++------- 2 files changed, 26 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 9c61510a6..128ba41ee 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index e87269a72..d49db49aa 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -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