Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 3 additions & 5 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -240,15 +227,13 @@ def __init__(
device: str,
enable_memory_saver: bool,
auxiliary_state_size: int,
owns_auxiliary_state_release: bool = False,
):
super().__init__(
size=size,
max_context_len=max_context_len,
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,
Expand Down Expand Up @@ -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()
Expand Down
4 changes: 0 additions & 4 deletions python/sglang/srt/hardware_backend/mlx/model_runner_stub.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
7 changes: 4 additions & 3 deletions python/sglang/srt/managers/schedule_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions python/sglang/srt/mem_cache/allocation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
10 changes: 1 addition & 9 deletions python/sglang/srt/mem_cache/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
16 changes: 13 additions & 3 deletions python/sglang/srt/mem_cache/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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 "
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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}"
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
Expand All @@ -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)

Expand Down
9 changes: 6 additions & 3 deletions test/registered/unit/layers/test_minicpm_sparse_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down
15 changes: 11 additions & 4 deletions test/registered/unit/managers/test_prefill_adder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
22 changes: 22 additions & 0 deletions test/registered/unit/mem_cache/test_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading