[session] fix mamba pool leak in StreamingSession.release_session + plumb idle leak check (#23496)

This commit is contained in:
Sam Shleifer
2026-05-02 11:38:08 +08:00
committed by GitHub
parent d41e8c459d
commit 63f225ca2e
4 changed files with 44 additions and 1 deletions
@@ -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:
@@ -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
@@ -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)
@@ -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: