diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 69c189b1bc88..1852c43ecac8 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1377,13 +1377,11 @@ def pop_preallocated( # Hybrid models (e.g. K3 with KDA): guard against prealloc # draining the mamba pool before the KV pool (would assert "Not # enough space for mamba cache"). Evict a cached mamba slot from - # the radix tree first (only if it manages mamba states; - # ChunkCache.evict is a no-op), else stop. + # the radix tree first (a no-op with the radix cache disabled), + # else stop. mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None) if mamba_allocator is not None and mamba_allocator.available_size() <= 0: - supports_mamba = self.tree_cache.supports_mamba() - if supports_mamba and hasattr(self.tree_cache, "evict"): - self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + self.tree_cache.evict(EvictParams(num_tokens=0, mamba_num=1)) if mamba_allocator.available_size() <= 0: break diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py index 9a45f49f246e..85bd3a19962b 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py @@ -216,21 +216,8 @@ def has_snapshot(self, index: Any) -> bool: class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool): - """Req-to-token pool with MLX auxiliary-state slot bookkeeping. - - Auxiliary-slot release has exactly one owner per configuration: - - * Radix cache enabled: the ``MlxAuxiliaryStateComponent`` of the unified - radix cache owns release — on finish it either frees the slot or - transfers it to the tree, nulling ``req.kv.mamba_pool_idx`` before the - request row is freed. The pool must NOT free auxiliary slots itself. - * Radix cache disabled (``ChunkCache``): no tree component exists, and - ``release_kv_cache``'s ``free_mamba_cache`` fallback is gated on - ``HybridReqToTokenPool``, which this pool is not — so the pool itself - owns release. Construct with ``owns_auxiliary_state_release=True`` and - ``free(req)`` returns the slot together with the request row; without - this, every finished request leaks its slot until allocation asserts. - """ + """Req-to-token pool with MLX auxiliary-state slot bookkeeping. It never + frees auxiliary slots: ``MlxAuxiliaryStateComponent`` owns their release.""" def __init__( self, @@ -240,7 +227,6 @@ def __init__( device: str, enable_memory_saver: bool, auxiliary_state_size: int, - owns_auxiliary_state_release: bool = False, ): super().__init__( size=size, @@ -248,7 +234,6 @@ def __init__( device=device, enable_memory_saver=enable_memory_saver, ) - self._owns_auxiliary_state_release = owns_auxiliary_state_release self.mamba_pool = MlxAuxiliaryStatePool( size=auxiliary_state_size, device=device, @@ -310,16 +295,6 @@ def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None): mamba_ping_pong_track_buffer_to_keep=track_buffer_to_keep, ) - def free(self, req): - if self._owns_auxiliary_state_release: - # No-radix configuration: nothing else will ever release the - # auxiliary slot, so return it with the request row. Keyed on - # req.kv.mamba_pool_idx (None-safe, nulled by free_mamba_cache), NOT - # on req_index_to_auxiliary_state_index_mapping, which may point - # at a slot the radix tree owns. - self.free_mamba_cache(req) - super().free(req) - def clear(self): super().clear() self.auxiliary_state_pool.clear() diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py b/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py index d5df916479e7..4bc3ade314f2 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py @@ -296,10 +296,6 @@ def initialize(self): device="cpu", enable_memory_saver=False, auxiliary_state_size=auxiliary_state_size, - # With the radix cache disabled no tree component exists to - # release auxiliary slots, so the pool owns their release - # (see MlxAuxiliaryStateReqToTokenPool docstring). - owns_auxiliary_state_release=get_memory().disable_radix_cache, ) else: self.req_to_token_pool = ReqToTokenPool( diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 53fca050c044..b508eacb13ff 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -652,7 +652,9 @@ def __init__( self.exact_chunk_fill = ( _use_exact_chunk_fill() and dllm_config is None - and not tree_cache.supports_mamba() + and not ( + tree_cache.supports_mamba() and tree_cache.supports_prefix_sharing() + ) ) if self.dllm_config is not None: @@ -695,8 +697,7 @@ def __init__( self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba() # A new state slot eats shared-gap bytes that `rem_total_tokens` counts # as free, so reserve per slot or admission over-commits. Gate on the - # ALLOCATOR, not `is_hybrid_ssm_cache`: that is False for `ChunkCache`, - # which would skip the reservation on the chunk-cache path. + # ALLOCATOR: the reservation holds with the radix cache disabled too. self._mamba_slot_cost = 0 if isinstance( self.token_to_kv_pool_allocator, diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 0001521e118d..8c0678c826ae 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -101,15 +101,9 @@ def _check_full_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str] session_held = self.pool_stats_observer.session_held_full_tokens() total = ps.full_capacity elif self.is_hybrid_ssm: - # Branch on cache type for the protected accessor (a mamba-capable - # cache splits full/mamba; ChunkCache only has the single protected_size). - # Use the allocator's `.size` for `total`: static max_total_num_tokens for - # non-unified pools, the dynamic byte-coordinated cap (matching - # `available_size`) for the unified pool. - if self.tree_cache.supports_mamba(): - protected = self.tree_cache.full_protected_size() - else: - protected = self.tree_cache.protected_size() + # `total` is the allocator's `.size`: static for non-unified pools, + # the byte-coordinated cap (matching `available_size`) for the unified pool. + protected = self.tree_cache.full_protected_size() session_held = self.pool_stats_observer.session_held_tokens() total = self.req_to_token_pool.schedulable_token_capacity( self.token_to_kv_pool_allocator.size diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index 542f8943f7b1..80aad30f1ed2 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -304,8 +304,8 @@ def alloc_req_slots( mamba_available_size = ( req_to_token_pool.mamba_allocator.schedulable_available_size() ) - # Eviction headroom factor: 3x (or lazy variant) for radix COW, 1x for chunk. - if tree_cache.supports_mamba(): + # Eviction headroom factor: 3x (or lazy variant) for radix COW, 1x without prefix sharing. + if tree_cache.supports_mamba() and tree_cache.supports_prefix_sharing(): factor = ( MAMBA_STATE_PER_REQ_PREFIX_CACHE_LAZY if req_to_token_pool.enable_mamba_extra_buffer_lazy diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 8a8fd5d360fe..6f82227b19c6 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -13,7 +13,7 @@ from sglang.srt.mem_cache.allocator.page_interleave import page_interleave_shard_size from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.hicache_storage import PoolTransfer -from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool +from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.unified_cache.component_type import ComponentType from sglang.srt.runtime_context import get_serving, get_spec from sglang.srt.utils.common import ceil_align @@ -334,14 +334,6 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr ) tree_cache.on_release(req, inserted=is_insert) - # If the prefix cache doesn't manage mamba states, we must free them here. - if isinstance(tree_cache.req_to_token_pool, HybridReqToTokenPool) and ( - not tree_cache.supports_mamba() - ): - assert req.kv.holds_mamba, ( - "mamba state is freed while the tree cache does not manage mamba states" - ) - tree_cache.req_to_token_pool.free_mamba_cache(req) # The DSV4-NPU ReqToTokenPool subclass's free() additionally releases the # c4/c128 state pages; other ReqToTokenPool subclasses are a no-op here. tree_cache.req_to_token_pool.free(req) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index c4d033f03154..638db7fb85ae 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -84,9 +84,12 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache: is_pure_swa = ctx.is_hybrid_swa and ctx.full_tokens_per_layer == 0 if ctx.disable_radix_cache and ( get_disagg().disaggregation_decode_retraction_backup == "host_pool" - # Streaming sessions live in UnifiedRadixCache; its disabled mode - # stands in for the chunk caches. Pure-SWA has no unified layout. - or (get_serving().enable_streaming_session and not is_pure_swa) + # Streaming sessions and mamba states need UnifiedRadixCache, whose + # disabled mode replaces the chunk caches; pure-SWA has no unified layout. + or ( + not is_pure_swa + and (get_serving().enable_streaming_session or ctx.is_hybrid_ssm) + ) ): return create_unified_radix_cache(ctx) @@ -292,6 +295,13 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: "a PR at https://github.com/sgl-project/sglang if you need this." ) + if ctx.is_hybrid_ssm and not cache.supports_mamba(): + raise NotImplementedError( + f"Models with mamba state are not verified with {type(cache).__name__}; " + "mamba state lives in UnifiedRadixCache. Please open an issue or a PR " + "at https://github.com/sgl-project/sglang if you need this." + ) + hicache_attached = cache.cache_controller is not None logger.info( "Tree cache initialized: source=%s impl=%s hybrid_swa=%s hybrid_ssm=%s " diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba.py index 57cac618d164..14e466917f3d 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba.py @@ -67,7 +67,9 @@ def __init__(self, cache: UnifiedRadixCache, params: CacheInitParams): assert isinstance(params.req_to_token_pool, HybridReqToTokenPool), ( f"MambaComponent requires HybridReqToTokenPool, got {type(params.req_to_token_pool)}" ) - if not params.enable_mamba_extra_buffer: + # Without the extra buffer only the sequence-end state exists, so a cached + # state needs page 1; a disabled tree caches none. + if not params.enable_mamba_extra_buffer and not params.disable: assert params.page_size == 1, ( f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}" ) diff --git a/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py b/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py index fdd1d43d0fb3..c7fde8309cfc 100644 --- a/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py +++ b/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py @@ -32,6 +32,9 @@ _SKIP_REASON = "requires mlx" if _HAS_MLX: + from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import ( + MlxAuxiliaryStateComponent, + ) from sglang.srt.hardware_backend.mlx.model_runner_stub import ( MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO as RATIO, ) @@ -265,15 +268,8 @@ def test_radix_disabled_default_sizing_is_one_to_one(self): self.assertEqual(stub.req_to_token_pool.auxiliary_state_pool.size, 3) def test_radix_disabled_sequential_requests_release_their_aux_slot(self): - # THE LEAK: with radix disabled, release_kv_cache's free_mamba_cache - # fallback never fires (the MLX pool is not a HybridReqToTokenPool) - # and ChunkCache frees token KV only, so pool.free(req) was the only - # release hook left -- and it freed just the request row. Every - # finished request permanently consumed one auxiliary slot and the - # (cap + 1)-th SEQUENTIAL request crashed with "Not enough MLX - # auxiliary state slots" even at concurrency 1. The pool now owns - # auxiliary release in this configuration: allocate/free/reallocate - # far past the pool size must succeed, with every slot returned. + # With radix disabled the component cleanup still frees the slot on finish; + # sequential requests far past the pool size must get every slot back. stub = _hybrid_stub_for_initialize( max_running_requests=2, max_mamba_cache_size=2, @@ -282,11 +278,16 @@ def test_radix_disabled_sequential_requests_release_their_aux_slot(self): with _arch(hybrid=True), _published(stub): stub.initialize() pool = stub.req_to_token_pool + component = MlxAuxiliaryStateComponent( + SimpleNamespace(req_to_token_pool=pool), + SimpleNamespace(enable_mamba_extra_buffer=False), + ) aux_capacity = pool.auxiliary_state_pool.available_size() for _ in range(3 * aux_capacity): req = _fake_req() self.assertIsNotNone(pool.alloc([req])) - pool.free(req) # as release_kv_cache does after ChunkCache + component.cleanup_after_caching_req(req=req, is_finished=True) + pool.free(req) self.assertIsNone(req.kv.mamba_pool_idx) self.assertEqual(pool.auxiliary_state_pool.available_size(), aux_capacity) diff --git a/test/registered/unit/layers/test_minicpm_sparse_cache.py b/test/registered/unit/layers/test_minicpm_sparse_cache.py index cf8d93b0e164..f15d94d034d7 100644 --- a/test/registered/unit/layers/test_minicpm_sparse_cache.py +++ b/test/registered/unit/layers/test_minicpm_sparse_cache.py @@ -188,8 +188,9 @@ def test_reserved_slots_are_excluded_from_full_pool_invariant(): swa_tokens_per_layer=None, max_total_num_tokens=64, tree_cache=SimpleNamespace( - supports_mamba=lambda: False, - protected_size=lambda: 0, + supports_mamba=lambda: True, + supports_prefix_sharing=lambda: False, + full_protected_size=lambda: 0, ), token_to_kv_pool_allocator=allocator, req_to_token_pool=pool, @@ -213,7 +214,9 @@ def test_hybrid_pool_stats_exclude_reserved_slots(): pool.mamba_allocator = SimpleNamespace(available_size=lambda: 1) pool.mamba_pool = SimpleNamespace(size=1) observer = SchedulerPoolStatsObserver( - tree_cache=SimpleNamespace(supports_mamba=lambda: False), + tree_cache=SimpleNamespace( + supports_mamba=lambda: True, supports_prefix_sharing=lambda: False + ), token_to_kv_pool_allocator=allocator, req_to_token_pool=pool, session_controller=None, diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 3a820900191e..85071759fd26 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -284,17 +284,24 @@ def test_continuation_without_limit_keeps_normal_chunk_size(self): def test_exact_chunk_fill_keeps_mamba_chunks_page_aligned(self): # A Mamba checkpoint only lands on a page-aligned chunk end, so an - # off-grid chunk leaves the rest of the prompt uncacheable. + # off-grid chunk leaves the rest of the prompt uncacheable. Without + # prefix sharing nothing is cached, so the chunk stays exact. self.mock_token_allocator.available_size.return_value = 32768 - self.mock_tree_cache.supports_prefix_sharing.return_value = False - for supports_mamba, expected in ((False, 100), (True, 64)): + cases = ((False, True, 100), (True, True, 64), (True, False, 100)) + for supports_mamba, supports_prefix_sharing, expected in cases: with ( - self.subTest(supports_mamba=supports_mamba), + self.subTest( + supports_mamba=supports_mamba, + supports_prefix_sharing=supports_prefix_sharing, + ), patch.object( schedule_policy, "_use_exact_chunk_fill", return_value=True ), ): self.mock_tree_cache.supports_mamba.return_value = supports_mamba + self.mock_tree_cache.supports_prefix_sharing.return_value = ( + supports_prefix_sharing + ) adder = self.create_adder( self.create_running_batch(), page_size=64, rem_chunk_tokens=100 ) diff --git a/test/registered/unit/mem_cache/test_registry.py b/test/registered/unit/mem_cache/test_registry.py index 3fec71a4c951..39d434be01b4 100644 --- a/test/registered/unit/mem_cache/test_registry.py +++ b/test/registered/unit/mem_cache/test_registry.py @@ -208,6 +208,28 @@ def test_streaming_with_disable_radix_keeps_pure_swa_off_unified(self): create_unified.assert_not_called() self.assertIs(result, PureSWAChunkCache.return_value) + def test_mamba_with_disable_radix_routes_to_unified(self): + ctx = _make_ctx( + self, + effective_chunked_prefill_size=512, + disable_radix_cache=True, + is_hybrid_ssm=True, + ) + with patch( + "sglang.srt.mem_cache.registry.create_unified_radix_cache" + ) as create_unified: + result = default_radix_cache_factory(ctx) + create_unified.assert_called_once_with(ctx) + self.assertIs(result, create_unified.return_value) + + def test_mamba_rejected_on_cache_without_mamba(self): + inner = MagicMock() + inner.supports_mamba.return_value = False + register_radix_cache_backend("nomamba", MagicMock(return_value=inner)) + + with self.assertRaisesRegex(NotImplementedError, "not verified"): + create_tree_cache(_make_ctx(self, backend="nomamba", is_hybrid_ssm=True)) + def test_unified_radix_cache_is_the_default(self): ctx = _make_ctx( self,