Add SWA support for runtime busy memory check (#21499)

This commit is contained in:
Liangsheng Yin
2026-04-12 00:39:51 -07:00
committed by GitHub
parent 19cb918653
commit f2377a00cb
@@ -8,7 +8,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.observability.metrics_collector import QueueCount
from sglang.srt.utils.common import ceil_align, raise_error_or_warn
@@ -323,33 +322,43 @@ class SchedulerRuntimeCheckerMixin:
)
return leak, msg
def _get_batch_uncached_size(self: Scheduler, batch: ScheduleBatch) -> int:
ret = 0
for req in batch.reqs:
assert req.kv_committed_freed == req.kv_overallocated_freed
uncached_len = 0
if not req.kv_committed_freed:
def _get_total_uncached_sizes(self: Scheduler) -> Tuple[int, int]:
"""Sum uncached tokens for full and SWA pools across all active batches.
Returns (full_uncached, swa_uncached). For non-SWA models, swa_uncached is 0.
For full pool: uncached = allocated - cache_protected_len
For SWA pool: uncached = allocated - max(cache_protected_len, swa_evicted_seqlen)
"""
# After decode: running_batch IS last_batch (same object), count once.
# After prefill: they differ, both hold uncached tokens.
batches = [self.last_batch]
if (
self.running_batch not in (None, self.last_batch)
and not self.running_batch.is_empty()
):
batches.append(self.running_batch)
full_uncached = 0
swa_uncached = 0
for batch in batches:
for req in batch.reqs:
assert req.kv_committed_freed == req.kv_overallocated_freed
if req.kv_committed_freed or req.req_pool_idx is None:
continue
allocated_len = req.kv_allocated_len
if self.page_size > 1:
allocated_len = ceil_align(allocated_len, self.page_size)
assert req.cache_protected_len % self.page_size == 0
uncached_len = allocated_len - req.cache_protected_len
ret += uncached_len
full_uncached += allocated_len - req.cache_protected_len
if self.is_hybrid_swa:
swa_uncached += allocated_len - max(
req.cache_protected_len, req.swa_evicted_seqlen
)
return ret
def _get_total_uncached_size(self: Scheduler) -> int:
"""Sum uncached tokens across the current and running batches."""
current_batch: ScheduleBatch = self.last_batch
uncached_size = self._get_batch_uncached_size(current_batch)
if (
current_batch.forward_mode.is_extend()
and self.running_batch is not None
and not self.running_batch.is_empty()
):
uncached_size += self._get_batch_uncached_size(self.running_batch)
return uncached_size
return full_uncached, swa_uncached
def self_check_during_busy(self: Scheduler):
if self.last_batch is None:
@@ -362,12 +371,21 @@ class SchedulerRuntimeCheckerMixin:
)
return
uncached = self._get_total_uncached_size()
leak, msg = self._check_full_pool(self.get_pool_stats(), uncached=uncached)
ps = self.get_pool_stats()
full_uncached, swa_uncached = self._get_total_uncached_sizes()
full_leak, full_msg = self._check_full_pool(ps, uncached=full_uncached)
swa_leak, swa_msg = False, ""
if self.is_hybrid_swa:
swa_leak, swa_msg = self._check_swa_pool(ps, uncached=swa_uncached)
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get() > 1:
logger.info(f"[Mem Check (BUSY)] {msg}")
assert not leak, f"Mem Leak Detected! {msg}"
logger.info(f"[Mem Check (BUSY)] {full_msg}")
if swa_msg:
logger.info(f"[Mem Check (BUSY)] {swa_msg}")
assert not full_leak, f"Full Pool Mem Leak Detected! {full_msg}"
assert not swa_leak, f"SWA Pool Mem Leak Detected! {swa_msg}"
def _check_req_pool(self: Scheduler):
if self.disaggregation_mode == DisaggregationMode.DECODE: