From 63f225ca2eb531de91f8efbbab700fe64d210378 Mon Sep 17 00:00:00 2001 From: Sam Shleifer Date: Fri, 1 May 2026 23:38:08 -0400 Subject: [PATCH] [session] fix mamba pool leak in StreamingSession.release_session + plumb idle leak check (#23496) --- .../scheduler_runtime_checker_mixin.py | 5 ++- .../sglang/srt/mem_cache/base_prefix_cache.py | 3 ++ .../srt/mem_cache/unified_radix_cache.py | 3 ++ .../sglang/srt/session/streaming_session.py | 34 +++++++++++++++++++ 4 files changed, 44 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index 72b01cfd8..47f9a4b31 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -166,6 +166,9 @@ class SchedulerRuntimeCheckerMixin: def _session_held_req_count(self: Scheduler) -> int: return self.tree_cache.session_held_req_count() + def _session_held_mamba_slots(self: Scheduler) -> int: + return self.tree_cache.session_held_mamba_slots(self._active_pool_idxs()) + def get_pool_stats(self: Scheduler) -> PoolStats: if self.is_hybrid_swa: pool_stats = self._get_swa_token_info() @@ -335,7 +338,7 @@ class SchedulerRuntimeCheckerMixin: ps.mamba_available_size, ps.mamba_evictable_size, self.tree_cache.mamba_protected_size(), - 0, + self._session_held_mamba_slots(), self.req_to_token_pool.mamba_pool.size, ) if leak: diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index 6c66da5fe..403dc6ed4 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -280,6 +280,9 @@ class BasePrefixCache(ABC, PrefixCacheTrait): def session_held_req_count(self, active_pool_idxs: Optional[set] = None) -> int: return 0 + def session_held_mamba_slots(self, active_pool_idxs: Optional[set] = None) -> int: + return 0 + def is_chunk_cache(self) -> bool: return False diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index bdaf77fc6..38dcb5652 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -812,6 +812,9 @@ class UnifiedRadixCache(BasePrefixCache): def session_held_req_count(self, active_pool_idxs: Optional[set] = None) -> int: return self.session.session_held_req_count(active_pool_idxs) + def session_held_mamba_slots(self, active_pool_idxs: Optional[set] = None) -> int: + return self.session.session_held_mamba_slots(active_pool_idxs) + def evictable_size(self) -> int: return self.component_evictable_size_.get(BASE_COMPONENT_TYPE, 0) diff --git a/python/sglang/srt/session/streaming_session.py b/python/sglang/srt/session/streaming_session.py index c52e9baee..a60b3376c 100644 --- a/python/sglang/srt/session/streaming_session.py +++ b/python/sglang/srt/session/streaming_session.py @@ -408,6 +408,8 @@ class StreamingSession(BasePrefixCache): self.token_to_kv_pool_allocator.free(kv_indices) self.req_to_token_pool.free_slots.append(slot.req_pool_idx) + self._free_slot_mamba(slot) + def session_held_tokens(self, active_pool_idxs: Optional[set] = None) -> int: """Total KV tokens held by session slots, not tracked by the tree. @@ -454,6 +456,38 @@ class StreamingSession(BasePrefixCache): return sum(_owned(s) for s in self.slots.values()) + def session_held_mamba_slots(self, active_pool_idxs: Optional[set] = None) -> int: + """Total mamba_pool entries held by session slots (mamba_pool_idx + + mamba_ping_pong_track_buffer). Excludes slots whose owning req is + currently in the batch -- those slots are counted via the normal + alloc/free paths (same convention as the sibling ``session_held_*`` + accessors). + """ + total = 0 + for slot in self.slots.values(): + in_batch = ( + active_pool_idxs is not None and slot.req_pool_idx in active_pool_idxs + ) + if in_batch: + continue + if slot.mamba_pool_idx is not None: + total += slot.mamba_pool_idx.numel() + if slot.mamba_ping_pong_track_buffer is not None: + total += slot.mamba_ping_pong_track_buffer.numel() + return total + + def _free_slot_mamba(self, slot: SessionSlot) -> None: + """Return a session slot's mamba pool state to the allocator.""" + mamba_pool = getattr(self.req_to_token_pool, "mamba_pool", None) + if mamba_pool is None: + return + if slot.mamba_pool_idx is not None: + mamba_pool.free(slot.mamba_pool_idx.unsqueeze(0)) + slot.mamba_pool_idx = None + if slot.mamba_ping_pong_track_buffer is not None: + mamba_pool.free(slot.mamba_ping_pong_track_buffer) + slot.mamba_ping_pong_track_buffer = None + # -- Internal helpers (streaming body bits) -- def _free_tail(self, slot: SessionSlot, req: Req, prefix_len: int) -> None: