From e57323cae910743185be763bb5c836dbb61f9996 Mon Sep 17 00:00:00 2001 From: shuwenn <47200617+alphabetc1@users.noreply.github.com> Date: Sun, 7 Jun 2026 10:46:29 +0800 Subject: [PATCH] [mem_cache][4/N] refactor: extract MambaTokenToKVPoolAllocator into allocator/ (#27256) --- python/sglang/srt/disaggregation/decode.py | 2 +- python/sglang/srt/managers/schedule_batch.py | 4 +- python/sglang/srt/managers/scheduler.py | 24 ++--- .../batch_result_processor.py | 2 +- .../scheduler_components/invariant_checker.py | 7 +- .../pool_stats_observer.py | 2 +- .../sglang/srt/mem_cache/allocator/mamba.py | 90 +++++++++++++++++++ python/sglang/srt/mem_cache/common.py | 4 +- .../srt/mem_cache/hi_mamba_radix_cache.py | 8 +- .../hybrid_cache/hybrid_pool_assembler.py | 3 + .../sglang/srt/mem_cache/mamba_radix_cache.py | 14 +-- python/sglang/srt/mem_cache/memory_pool.py | 72 +++------------ .../sglang/srt/mem_cache/memory_pool_host.py | 4 +- .../mamba_component.py | 18 ++-- .../srt/mem_cache/unified_radix_cache.py | 2 +- .../sglang/srt/session/streaming_session.py | 8 +- .../unit/mem_cache/test_mamba_unittest.py | 31 ++++--- .../test_unified_radix_cache_unittest.py | 2 +- 18 files changed, 177 insertions(+), 120 deletions(-) create mode 100644 python/sglang/srt/mem_cache/allocator/mamba.py diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index b5ede8691..34ba4e0df 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -241,7 +241,7 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): def clear(self): self.free_slots = list(range(1, self._alloc_size)) - self.mamba_pool.clear() + self.mamba_allocator.clear() @dataclass diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index cc24f773b..182aeeb46 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2417,10 +2417,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): if envs.SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL.get(): new_slot = None else: - new_slot = pool.mamba_pool.alloc(1) + new_slot = pool.mamba_allocator.alloc(1) if new_slot is None: self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1)) - new_slot = pool.mamba_pool.alloc(1) + new_slot = pool.mamba_allocator.alloc(1) if new_slot is not None: pool.set_mamba_ping_pong_slot(req, other_idx, new_slot[0]) req.mamba_next_track_idx = other_idx diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index eda6cae0e..da66e6cff 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2662,9 +2662,9 @@ class Scheduler( self.running_batch.reqs, ) - mamba_pool = getattr(self.req_to_token_pool, "mamba_pool", None) - if mamba_pool is not None: - mamba_pool.alloc_group_begin(len(self.waiting_queue)) + mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None) + if mamba_allocator is not None: + mamba_allocator.alloc_group_begin(len(self.waiting_queue)) # Get requests from the waiting queue to a new prefill batch for req in self.waiting_queue: if self.enable_lora and not self._can_schedule_lora_req(req, running_loras): @@ -2725,14 +2725,14 @@ class Scheduler( and req.mamba_pool_idx is not None and not getattr(req, "session", None) ): - self.tree_cache.req_to_token_pool.mamba_pool.free( + self.tree_cache.req_to_token_pool.mamba_allocator.free( req.mamba_pool_idx.unsqueeze(-1) ) req.mamba_pool_idx = None break - if mamba_pool is not None: - mamba_pool.alloc_group_end() + if mamba_allocator is not None: + mamba_allocator.alloc_group_end() # Update waiting queue can_run_list: List[Req] = adder.can_run_list @@ -2867,9 +2867,13 @@ class Scheduler( ): old_available_tokens = self.token_to_kv_pool_allocator.available_size() old_ratio = self.new_token_ratio_tracker.current - mamba_pool = getattr(self.tree_cache.req_to_token_pool, "mamba_pool", None) + mamba_allocator = getattr( + self.tree_cache.req_to_token_pool, "mamba_allocator", None + ) old_mamba_available = ( - mamba_pool.available_size() if mamba_pool is not None else None + mamba_allocator.available_size() + if mamba_allocator is not None + else None ) retracted_reqs, new_token_ratio, reqs_to_abort = batch.retract_decode( self.server_args @@ -2877,8 +2881,8 @@ class Scheduler( new_available_tokens = self.token_to_kv_pool_allocator.available_size() new_token_gained = new_available_tokens - old_available_tokens mamba_num_gained = ( - mamba_pool.available_size() - old_mamba_available - if mamba_pool is not None + mamba_allocator.available_size() - old_mamba_available + if mamba_allocator is not None else None ) diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 8fa300afe..b0da5cd55 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -938,7 +938,7 @@ class SchedulerBatchResultProcessor: other_val = req.mamba_ping_pong_track_buffer[other_idx].item() if other_val != -1: pool = batch.req_to_token_pool - pool.mamba_pool.free( + pool.mamba_allocator.free( req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0) ) pool.set_mamba_ping_pong_slot(req, other_idx, -1) diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 6cf379d8c..4857f5498 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -135,13 +135,12 @@ class SchedulerInvariantChecker: leaked_full_pages = ( expected_full_pages - free_full_pages - cached_full_pages ) - free_mamba_pages = set( - self.req_to_token_pool.mamba_pool.free_slots.tolist() - ) + mamba_allocator = self.req_to_token_pool.mamba_allocator + free_mamba_pages = set(mamba_allocator.free_slots.tolist()) cached_mamba_pages = set( self.tree_cache.all_mamba_values_flatten().tolist() ) - expected_mamba_pages = set(range(self.req_to_token_pool.mamba_pool.size)) + expected_mamba_pages = set(range(1, mamba_allocator.size + 1)) leaked_mamba_pages = ( expected_mamba_pages - free_mamba_pages - cached_mamba_pages ) diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index a6ed75228..b700bd525 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -246,7 +246,7 @@ class SchedulerPoolStatsObserver: full_evictable_size = ( self.tree_cache.full_evictable_size() if is_mamba_radix_cache else 0 ) - mamba_available_size = self.req_to_token_pool.mamba_pool.available_size() + mamba_available_size = self.req_to_token_pool.mamba_allocator.available_size() mamba_evictable_size = ( self.tree_cache.mamba_evictable_size() if is_mamba_radix_cache else 0 ) diff --git a/python/sglang/srt/mem_cache/allocator/mamba.py b/python/sglang/srt/mem_cache/allocator/mamba.py new file mode 100644 index 000000000..b6ddbdf4a --- /dev/null +++ b/python/sglang/srt/mem_cache/allocator/mamba.py @@ -0,0 +1,90 @@ +""" +Copyright 2026 SGLang Team +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. + +Slot allocator for the Mamba state pool. + +Mamba caches one whole state tensor per request, so the allocator hands out +fixed-size slots (1 per request) rather than paged token KV indices. The +underlying tensor storage lives in ``MambaPool``; this class owns only the +free-slot bookkeeping. +""" + +from __future__ import annotations + +from typing import Iterator, Optional + +import torch + + +class MambaSlotAllocator: + """Manages the free-list of Mamba pool slot indices. + + Unlike ``BaseTokenToKVPoolAllocator`` which is designed for per-token KV + pages, Mamba slots are request-level (typically 1 slot per request). + We keep the interface minimal and do NOT inherit the KV base class. + """ + + def __init__(self, size: int, device: str): + self.size = size + self.device = device + # Active preallocated batch for `alloc_group_begin` / `alloc_group_end`. + # When non-None, `alloc(1)` consumes the next slot from this iterator + # instead of calling `_do_alloc(1)` per request. Reset to None outside + # a group window so `alloc` falls through to the per-call path. + self._alloc_iter: Optional[Iterator] = None + self.clear() + + def available_size(self) -> int: + return len(self.free_slots) + + def alloc_group_begin(self, num_reqs: int): + """Pre-allocate a batch of slots for match_prefix to amortize overhead.""" + self._alloc_iter = None + if num_reqs > 0: + result = self._do_alloc(num_reqs) + if result is not None: + self._alloc_iter = iter(result.split(1)) + + def alloc_group_end(self): + """Return any unused pre-allocated slots from the current group.""" + if self._alloc_iter is not None: + remaining = list(self._alloc_iter) + if remaining: + self.free(torch.cat(remaining)) + self._alloc_iter = None + + def alloc(self, need_size: int) -> Optional[torch.Tensor]: + if self._alloc_iter is not None and need_size == 1: + slot = next(self._alloc_iter, None) + if slot is not None: + return slot + return self._do_alloc(need_size) + + def _do_alloc(self, need_size: int) -> Optional[torch.Tensor]: + if need_size > len(self.free_slots): + return None + select_index = self.free_slots[:need_size] + self.free_slots = self.free_slots[need_size:] + return select_index + + def free(self, free_index: torch.Tensor): + if free_index.numel() == 0: + return + self.free_slots = torch.cat((self.free_slots, free_index)) + + def clear(self): + # Slot 0 is reserved as a dummy write target for padded tokens. + self.free_slots = torch.arange( + 1, self.size + 1, dtype=torch.int64, device=self.device + ) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 2139e2902..f201c2257 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -314,7 +314,7 @@ def alloc_req_slots( """Allocate request slots from the pool.""" num_reqs = len(reqs) if isinstance(req_to_token_pool, HybridReqToTokenPool): - mamba_available_size = req_to_token_pool.mamba_pool.available_size() + mamba_available_size = req_to_token_pool.mamba_allocator.available_size() if tree_cache.supports_mamba(): factor = ( MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY @@ -486,7 +486,7 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr ), "Only MambaRadixCache allow freeing before alloc" # TODO (csy, hanming): clean up this early allocation logic if req.mamba_pool_idx is not None: - tree_cache.req_to_token_pool.mamba_pool.free( + tree_cache.req_to_token_pool.mamba_allocator.free( req.mamba_pool_idx.unsqueeze(-1) ) req.mamba_pool_idx = None diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index b028574b6..47818e930 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -508,7 +508,7 @@ class HiMambaRadixCache(MambaRadixCache): if node.mamba_value is None: return 0 mamba_num = len(node.mamba_value) - self.req_to_token_pool.mamba_pool.free(node.mamba_value) + self.req_to_token_pool.mamba_allocator.free(node.mamba_value) if node.mamba_lock_ref > 0: self.mamba_protected_size_ -= mamba_num node.mamba_lock_ref = 0 @@ -797,7 +797,7 @@ class HiMambaRadixCache(MambaRadixCache): # Internal: free device mamba only, KV stays on device (tombstone) x_next = self.mamba_lru_list.get_prev_no_lock(x) mamba_num_evicted += len(x.mamba_value) - self.req_to_token_pool.mamba_pool.free(x.mamba_value) + self.req_to_token_pool.mamba_allocator.free(x.mamba_value) self.mamba_lru_list.remove_node(x) self._tombstone_internal_node(x) else: @@ -1047,7 +1047,7 @@ class HiMambaRadixCache(MambaRadixCache): if cow_mamba and mamba_node.mamba_value is not None: if req.mamba_pool_idx is None: dst_index = self._alloc_with_evict( - self.req_to_token_pool.mamba_pool, + self.req_to_token_pool.mamba_allocator, 1, self.evict_mamba, lock_node=mamba_node, @@ -2085,7 +2085,7 @@ class HiMambaRadixCache(MambaRadixCache): ): if req.mamba_pool_idx is None: req.mamba_pool_idx = self._alloc_with_evict( - self.req_to_token_pool.mamba_pool, + self.req_to_token_pool.mamba_allocator, len(last_hit_node.mamba_host_value), self.evict_mamba, lock_node=last_hit_node, diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index bdf4611c7..a76256639 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -515,6 +515,7 @@ def build_hybrid_mamba_stack( enable_storage_metrics: bool = False, ) -> tuple[HostPoolGroup, HybridCacheController]: transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping) + mamba_allocator = params.req_to_token_pool.mamba_allocator kv_host_pool = build_kv_host_pool( kv_pool=kv_pool, page_size=page_size, @@ -545,6 +546,8 @@ def build_hybrid_mamba_stack( transfer_layer_num=transfer_layer_num, host_evict_fn=host_mamba_evict_fn, device_evict_fn=device_mamba_evict_fn, + device_alloc_fn=mamba_allocator.alloc, + device_free_fn=mamba_allocator.free, ), ] host_pool_group = HostPoolGroup(entries) diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index 36ee19eab..59903afe5 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -671,7 +671,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) new_prefix_len, mamba_exist = result.prefix_len, result.mamba_exist if mamba_exist: - self.req_to_token_pool.mamba_pool.free(mamba_value_donated) + self.req_to_token_pool.mamba_allocator.free(mamba_value_donated) # The prefix indices could be updated, reuse it match_result = self.match_prefix( @@ -729,7 +729,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): self._record_remove_event(x) self.token_to_kv_pool_allocator.free(x.value) full_num_evicted = len(x.value) - self.req_to_token_pool.mamba_pool.free(x.mamba_value) + self.req_to_token_pool.mamba_allocator.free(x.mamba_value) mamba_num_evicted = len(x.mamba_value) # 2. get the next node, update the lru lists @@ -782,7 +782,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): if len(x.children) > 0: # 1. an internal node, free mamba tokens. - self.req_to_token_pool.mamba_pool.free(x.mamba_value) + self.req_to_token_pool.mamba_allocator.free(x.mamba_value) mamba_num_evicted += len(x.mamba_value) # 2. get the next node, update the lru lists @@ -947,10 +947,10 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): def _alloc_mamba_slot(self) -> torch.Tensor: """Allocate one mamba pool slot, evicting if necessary.""" - slot = self.req_to_token_pool.mamba_pool.alloc(1) + slot = self.req_to_token_pool.mamba_allocator.alloc(1) if slot is None: self.evict(EvictParams(num_tokens=0, mamba_num=1)) - slot = self.req_to_token_pool.mamba_pool.alloc(1) + slot = self.req_to_token_pool.mamba_allocator.alloc(1) assert slot is not None, "Can not alloc mamba cache" return slot @@ -1046,11 +1046,11 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): # Defer COW to forward stream: record source index, allocate destination if cow_mamba and last_node.mamba_value is not None: if req.mamba_pool_idx is None: - dst_index = self.req_to_token_pool.mamba_pool.alloc(1) + dst_index = self.req_to_token_pool.mamba_allocator.alloc(1) if dst_index is None: self.inc_lock_ref(last_node) self.evict(EvictParams(num_tokens=0, mamba_num=1)) - dst_index = self.req_to_token_pool.mamba_pool.alloc(1) + dst_index = self.req_to_token_pool.mamba_allocator.alloc(1) self.dec_lock_ref(last_node) assert dst_index is not None, "Can not alloc mamba cache" req.mamba_pool_idx = dst_index[0] diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index bfcb1e9a1..9332fa493 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -27,7 +27,7 @@ import dataclasses import logging from contextlib import contextmanager, nullcontext from dataclasses import dataclass, fields -from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union import numpy as np import torch @@ -44,6 +44,7 @@ from sglang.srt.layers.attention.dsa.quant_k_cache import ( from sglang.srt.layers.attention.dsa.utils import aiter_can_use_preshuffle_paged_mqa from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz from sglang.srt.layers.radix_attention import RadixAttention +from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator from sglang.srt.mem_cache.triton_ops.cache_move import ( copy_all_layer_kv_cache_tiled, ) @@ -360,19 +361,9 @@ class MambaPool: f"conv_state size: {get_tensor_size_bytes(conv_state) / GB:.2f}GB, " f"ssm_state size: {get_tensor_size_bytes(temporal_state) / GB:.2f}GB " ) - # The padded slot 0 is used for writing dummy outputs from padded tokens. - self.free_slots = torch.arange( - 1, self.size + 1, dtype=torch.int64, device=self.device - ) self.mem_usage = self.mamba_cache.mem_usage_bytes() / GB self.num_mamba_layers = num_mamba_layers - # Active preallocated batch for `alloc_group_begin` / `alloc_group_end`. - # When non-None, `alloc(1)` consumes the next slot from this iterator - # instead of calling `_do_alloc(1)` per request. Reset to None outside - # a group window so `alloc` falls through to the per-call path. - self._alloc_iter: Optional[Iterator] = None - def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState: assert isinstance(self.mamba_cache, self.SpeculativeState) return self.mamba_cache @@ -380,39 +371,6 @@ class MambaPool: def mamba2_layer_cache(self, layer_id: int): return self.mamba_cache.at_layer_idx(layer_id) - def available_size(self): - return len(self.free_slots) - - # -- Batched alloc for match_prefix -- - def alloc_group_begin(self, num_reqs: int): - self._alloc_iter = None - if num_reqs > 0: - result = self._do_alloc(num_reqs) - if result is not None: - self._alloc_iter = iter(result.split(1)) - - def alloc_group_end(self): - if self._alloc_iter is not None: - remaining = list(self._alloc_iter) - if remaining: - self.free(torch.cat(remaining)) - self._alloc_iter = None - - def alloc(self, need_size: int) -> Optional[torch.Tensor]: - if self._alloc_iter is not None and need_size == 1: - slot = next(self._alloc_iter, None) - if slot is not None: - return slot - return self._do_alloc(need_size) - - def _do_alloc(self, need_size: int) -> Optional[torch.Tensor]: - if need_size > len(self.free_slots): - return None - - select_index = self.free_slots[:need_size] - self.free_slots = self.free_slots[need_size:] - return select_index - def clear_slots(self, indices: torch.Tensor): """Zero out mamba state at the given pool indices. Must run on forward stream.""" need_size = len(indices) @@ -428,16 +386,6 @@ class MambaPool: ) t[:, indices] = z - def free(self, free_index: torch.Tensor): - if free_index.numel() == 0: - return - self.free_slots = torch.cat((self.free_slots, free_index)) - - def clear(self): - self.free_slots = torch.arange( - 1, self.size + 1, dtype=torch.int64, device=self.device - ) - def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor): for i in range(len(self.mamba_cache.conv)): self.mamba_cache.conv[i][:, dst_indices] = self.mamba_cache.conv[i][ @@ -588,6 +536,10 @@ class HybridReqToTokenPool(ReqToTokenPool): enable_memory_saver=self.enable_memory_saver, speculative_num_draft_tokens=speculative_num_draft_tokens, ) + self.mamba_allocator = MambaSlotAllocator( + size=mamba_size, + device=device, + ) self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)} self.device = device @@ -622,10 +574,10 @@ class HybridReqToTokenPool(ReqToTokenPool): if req.mamba_pool_idx is not None: # for radix cache / continuing chunked pass else: - mid = self.mamba_pool.alloc(1) + mid = self.mamba_allocator.alloc(1) assert ( mid is not None - ), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_pool.available_size()=}, {len(reqs)=}" + ), f"Not enough space for mamba cache, try to increase --mamba-full-memory-ratio or --max-mamba-cache-size. {mid=}, {self.mamba_pool.size=}, {self.mamba_allocator.available_size()=}, {len(reqs)=}" req.mamba_pool_idx = mid[0] req.mamba_needs_clear = True mamba_indices.append(req.mamba_pool_idx) @@ -694,7 +646,7 @@ class HybridReqToTokenPool(ReqToTokenPool): if self.enable_mamba_extra_buffer_lazy else self.mamba_ping_pong_track_buffer_size ) - slots = self.mamba_pool.alloc(n) + slots = self.mamba_allocator.alloc(n) assert slots is not None, ( "Not enough space for mamba ping pong idx, " "try to increase --mamba-full-memory-ratio." @@ -749,7 +701,7 @@ class HybridReqToTokenPool(ReqToTokenPool): ): mamba_index = req.mamba_pool_idx assert mamba_index is not None, "double free? mamba_index is None" - self.mamba_pool.free(mamba_index.unsqueeze(0)) + self.mamba_allocator.free(mamba_index.unsqueeze(0)) req.mamba_pool_idx = None if self.enable_mamba_extra_buffer: @@ -789,7 +741,7 @@ class HybridReqToTokenPool(ReqToTokenPool): mamba_ping_pong_track_buffer_to_free != -1 ] ) - self.mamba_pool.free(mamba_ping_pong_track_buffer_to_free) + self.mamba_allocator.free(mamba_ping_pong_track_buffer_to_free) # Match the req.mamba_pool_idx=None clear above so the next # alloc() doesn't see a stale ping-pong reference on the req # and skip allocation (which would silently reuse a freed @@ -800,7 +752,7 @@ class HybridReqToTokenPool(ReqToTokenPool): def clear(self): logger.info("Reset HybridReqToTokenPool") super().clear() - self.mamba_pool.clear() + self.mamba_allocator.clear() self.req_index_to_mamba_index_mapping.zero_() if self.enable_mamba_extra_buffer: self.req_index_to_mamba_ping_pong_track_buffer_mapping.zero_() diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 92e5e47b8..b42ad7ba6 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -2574,8 +2574,8 @@ class PoolEntry: device_evict_fn: Optional[Callable] = None # Optional alloc/free overrides for the device side, used by # _resolve_pool_transfers_allocation. Set when entry.device_pool is the - # raw KV pool (layout) rather than an allocator (e.g. SWA, where alloc - # lives on a separate sub-allocator inside SWATokenToKVPoolAllocator). + # raw KV/state pool (layout) rather than an allocator (e.g. SWA/Mamba, + # where alloc lives on a separate allocator object). # When None, fall back to entry.device_pool.alloc/free. device_alloc_fn: Optional[Callable] = None device_free_fn: Optional[Callable] = None diff --git a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py index d85fbef4d..5b1758bbd 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py @@ -96,7 +96,7 @@ class MambaComponent(TreeComponent): if cow_mamba and mamba_value is not None: assert req is not None if req.mamba_pool_idx is None: - dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1) + dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1) if dst_index is None: # Capture the inc result and thread swa_uuid_for_lock back # into dec. Without it, SWA's release walks past this @@ -105,7 +105,7 @@ class MambaComponent(TreeComponent): # on ancestor nodes. lock_result = self.cache.inc_lock_ref(last_node) self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) - dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1) + dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1) self.cache.dec_lock_ref(last_node, lock_result.to_dec_params()) assert dst_index is not None, "Can not alloc mamba cache" req.mamba_pool_idx = dst_index[0] @@ -172,7 +172,7 @@ class MambaComponent(TreeComponent): # Device layer if EvictLayer.DEVICE in target and cd.value is not None: - self.cache.req_to_token_pool.mamba_pool.free(cd.value) + self.cache.req_to_token_pool.mamba_allocator.free(cd.value) freed = len(cd.value) self.cache.component_evictable_size_[self.component_type] -= freed cd.value = None @@ -285,10 +285,10 @@ class MambaComponent(TreeComponent): def _alloc_mamba_slot(self) -> torch.Tensor: """Allocate one mamba pool slot, evicting if necessary.""" - slot = self.cache.req_to_token_pool.mamba_pool.alloc(1) + slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1) if slot is None: self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) - slot = self.cache.req_to_token_pool.mamba_pool.alloc(1) + slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1) assert slot is not None, "Can not alloc mamba cache" return slot @@ -365,7 +365,9 @@ class MambaComponent(TreeComponent): if insert_params.mamba_value is not None and ( insert_result is None or insert_result.mamba_exist ): - self.cache.req_to_token_pool.mamba_pool.free(insert_params.mamba_value) + self.cache.req_to_token_pool.mamba_allocator.free( + insert_params.mamba_value + ) req.mamba_last_track_seqlen = None # ---- HiCache Hooks ---- @@ -414,10 +416,10 @@ class MambaComponent(TreeComponent): cd = node.component_data[ct] if req is not None and cd.host_value is not None: if req.mamba_pool_idx is None: - dst = self.cache.req_to_token_pool.mamba_pool.alloc(1) + dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1) if dst is None: self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) - dst = self.cache.req_to_token_pool.mamba_pool.alloc(1) + dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1) assert dst is not None, "Cannot alloc mamba for load_back" req.mamba_pool_idx = dst[0] transfers.append( diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index d31cda524..75f8b6d4f 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2511,7 +2511,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): if ct.is_swa: available_size = self.token_to_kv_pool_allocator.swa_available_size() elif ct.is_mamba: - available_size = self.req_to_token_pool.mamba_pool.available_size() + available_size = self.req_to_token_pool.mamba_allocator.available_size() else: continue diff --git a/python/sglang/srt/session/streaming_session.py b/python/sglang/srt/session/streaming_session.py index 475e9cf83..e1b0b18ef 100644 --- a/python/sglang/srt/session/streaming_session.py +++ b/python/sglang/srt/session/streaming_session.py @@ -500,14 +500,14 @@ class StreamingSession(BasePrefixCache): 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: + mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None) + if mamba_allocator is None: return if slot.mamba_pool_idx is not None: - mamba_pool.free(slot.mamba_pool_idx.unsqueeze(0)) + mamba_allocator.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) + mamba_allocator.free(slot.mamba_ping_pong_track_buffer) slot.mamba_ping_pong_track_buffer = None # -- Internal helpers (streaming body bits) -- diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index 315fb69d3..5dda30a73 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -109,7 +109,7 @@ class TestMamba(unittest.TestCase): ) assert req_to_token_pool.available_size() == max_num_reqs - assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size + assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size sampling_params = SamplingParams( temperature=0, @@ -125,34 +125,41 @@ class TestMamba(unittest.TestCase): # alloc req req_to_token_pool.alloc([req]) assert req_to_token_pool.available_size() == max_num_reqs - 1 - assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1 + assert ( + req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1 + ) # free req req_to_token_pool.free_mamba_cache(req) req_to_token_pool.free(req) assert req_to_token_pool.available_size() == max_num_reqs - assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size + assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size # alloc req without free mamba cache req.mamba_pool_idx = None req_to_token_pool.alloc([req]) req_to_token_pool.free(req) assert req_to_token_pool.available_size() == max_num_reqs - assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1 + assert ( + req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1 + ) # alloc again req_to_token_pool.alloc([req]) assert req_to_token_pool.available_size() == max_num_reqs - 1 - assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1 + assert ( + req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1 + ) def test_mamba_radix_cache_1(self): tree, allocator, req_to_token_pool, make_dummy_req = ( self._setup_tree_and_allocator() ) + mamba_allocator = req_to_token_pool.mamba_allocator mamba_pool = req_to_token_pool.mamba_pool # test print( - f"[Start] allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}" + f"[Start] allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}" ) req1 = make_dummy_req() req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3) @@ -170,7 +177,7 @@ class TestMamba(unittest.TestCase): ) prefix_len = result.prefix_len print( - f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}" + f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}" ) req2 = make_dummy_req() req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7) @@ -188,7 +195,7 @@ class TestMamba(unittest.TestCase): ) prefix_len = result.prefix_len print( - f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}" + f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}" ) req3 = make_dummy_req() @@ -207,7 +214,7 @@ class TestMamba(unittest.TestCase): ) prefix_len = result.prefix_len print( - f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}" + f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}" ) req4 = make_dummy_req() req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7) @@ -225,7 +232,7 @@ class TestMamba(unittest.TestCase): ) prefix_len = result.prefix_len print( - f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}" + f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}" ) tree.pretty_print() @@ -553,7 +560,7 @@ class TestMamba(unittest.TestCase): _, _, req_to_token_pool, _ = self._setup_tree_and_allocator() mamba_pool = req_to_token_pool.mamba_pool n = 3 - indices = mamba_pool.alloc(n) + indices = req_to_token_pool.mamba_allocator.alloc(n) self.assertIsNotNone(indices) # Write known sentinel values at the allocated slots. @@ -608,7 +615,7 @@ class TestMamba(unittest.TestCase): n_tokens = 4 kv_indices = allocator.alloc(n_tokens) self.assertIsNotNone(kv_indices) - mamba_indices = mamba_pool.alloc(1) + mamba_indices = req_to_token_pool.mamba_allocator.alloc(1) self.assertIsNotNone(mamba_indices) # Write sentinel values into KV buffers (all full-attention layers). diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 0bd21c648..4c35c0174 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -3545,7 +3545,7 @@ class UnifiedRadixCacheSuite: xfer = tree.components[ComponentType.MAMBA].build_hicache_transfers( node, CacheTransferPhase.LOAD_BACK )[0] - new_mamba = req_to_token_pool.mamba_pool.alloc(1) + new_mamba = req_to_token_pool.mamba_allocator.alloc(1) self.assertIsNotNone(new_mamba) xfer.device_indices = new_mamba tree.components[ComponentType.MAMBA].commit_hicache_transfer(