[session] fix mamba pool leak in StreamingSession.release_session + plumb idle leak check (#23496)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user