diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 62a39ba74..7d2c48985 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -3386,6 +3386,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): req_to_token_pool=self.req_to_token_pool, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, is_chunk_cache=self.tree_cache.is_chunk_cache(), + retain_floor=self.tree_cache.swa_retain_floor(req), ) def __str__(self): diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index 1f40409a4..0c24cb2f1 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -377,6 +377,13 @@ class BasePrefixCache(ABC, PrefixCacheTrait): def supports_swa(self) -> bool: return False + def swa_retain_floor(self, req) -> int | None: + # A match lands on a state checkpoint rather than on the tail, so a cache + # that pairs SWA with mamba/conv checkpoints has to keep the window behind + # the last checkpoint. Those caches override this. Everyone else has + # nothing deeper than the tail to protect. + return None + def swa_reprefill_tail_tokens(self) -> int: # Only the unified_kv compress-only HiCache layout needs to hold back a # trailing sliding window for re-prefill; every other cache keeps SWA diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 9d18b37e6..19e1c3e12 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -53,6 +53,7 @@ def free_swa_out_of_window_slots( req_to_token_pool: ReqToTokenPool, token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator, is_chunk_cache: bool = False, + retain_floor: int | None = None, ) -> None: if req.kv is None: return @@ -76,6 +77,12 @@ def free_swa_out_of_window_slots( # boundary (page_floor(seq_len)) so the last leaf is never all-tombstone. # No extra page margin is needed. evict_threshold = pre_len - max(sliding_window_size, page_size) + if retain_floor is not None and not is_chunk_cache: + # The caller owns where the floor is (see BasePrefixCache.swa_retain_floor); + # this only promises not to free past it. Chunk cache has no tree, so a + # retained checkpoint could never be matched and holding it is pure cost. + evict_threshold = min(evict_threshold, retain_floor) + new_swa_evicted_seqlen = max( req.kv.swa_evicted_seqlen, evict_threshold, diff --git a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py index 1074c7cbf..1ff56aaf1 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py @@ -755,6 +755,7 @@ class SWAComponent(TreeComponent): page_size=self.cache.page_size, req_to_token_pool=self.cache.req_to_token_pool, token_to_kv_pool_allocator=self.cache.token_to_kv_pool_allocator, + retain_floor=self.cache.swa_retain_floor(req), ) insert_params.swa_evicted_seqlen = req.kv.swa_evicted_seqlen diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 7483d0d8f..7e8669bf9 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2098,6 +2098,14 @@ class UnifiedRadixCache(BasePrefixCache): ) return swa.sliding_window_size if unified_compress_only_hicache else 0 + def swa_retain_floor(self, req) -> int | None: + if not self.is_mamba_enabled or self._sliding_window_size is None: + return None + checkpoint = req.mamba_last_track_seqlen + if checkpoint is None: + return None + return checkpoint - self._sliding_window_size + def supports_swa(self) -> bool: return self.is_swa_enabled diff --git a/test/registered/unit/mem_cache/test_swa_eviction_boundary.py b/test/registered/unit/mem_cache/test_swa_eviction_boundary.py index a884f0e44..7ab78ffe0 100644 --- a/test/registered/unit/mem_cache/test_swa_eviction_boundary.py +++ b/test/registered/unit/mem_cache/test_swa_eviction_boundary.py @@ -21,6 +21,7 @@ import torch from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.cache_init_params import CacheInitParams +from sglang.srt.mem_cache.common import free_swa_out_of_window_slots from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache @@ -196,6 +197,146 @@ class TestSWAEvictionBoundary(unittest.TestCase): ) tree.sanity_check() + # -- Retention floor: never free past the last state checkpoint -- + + def test_retain_floor_clamps_eviction(self): + """A hybrid cache keeps SWA down to the last state checkpoint, not to the + window behind the tail, because that is where a prefix match lands. The + floor must clamp the frontier even though the tail has moved far past it.""" + page_size, window = 8, 16 + tree, allocator, pool = _build_swa_tree( + page_size=page_size, sliding_window_size=window + ) + seq_len = 200 + checkpoint = 96 + kv = _swa_alloc(allocator, seq_len) + pool.write((0, slice(0, seq_len)), kv) + req = _make_req(0, list(range(seq_len)), 0, tree) + batch = _make_batch(tree, allocator, pool) + + free_swa_out_of_window_slots( + req, + seq_len - 1, + sliding_window_size=window, + page_size=page_size, + req_to_token_pool=batch.req_to_token_pool, + token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator, + retain_floor=checkpoint - window, + ) + + # Without the floor this would reach page_floor(199 - 16) = 176. + self.assertLessEqual(req.kv.swa_evicted_seqlen, checkpoint - window) + self.assertEqual(req.kv.swa_evicted_seqlen % page_size, 0) + + def test_retain_floor_ignored_for_chunk_cache(self): + """Chunk cache builds no tree, so a retained checkpoint could never be + matched. Holding it would cost SWA slots for nothing.""" + page_size, window = 8, 16 + seq_len = 200 + tree, allocator, pool = _build_swa_tree( + page_size=page_size, sliding_window_size=window + ) + kv = _swa_alloc(allocator, seq_len) + pool.write((0, slice(0, seq_len)), kv) + req = _make_req(0, list(range(seq_len)), 0, tree) + batch = _make_batch(tree, allocator, pool) + + free_swa_out_of_window_slots( + req, + seq_len - 1, + sliding_window_size=window, + page_size=page_size, + req_to_token_pool=batch.req_to_token_pool, + token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator, + is_chunk_cache=True, + retain_floor=16, + ) + + expected = (seq_len - 1 - window) // page_size * page_size + self.assertEqual(req.kv.swa_evicted_seqlen, expected) + + def test_retain_floor_none_matches_old_behaviour(self): + """retain_floor=None must reproduce the pre-change frontier exactly, so a + cache without a second state stream is unaffected.""" + page_size, window = 8, 16 + seq_len = 200 + frontiers = [] + for floor in (None, "absent"): + tree, allocator, pool = _build_swa_tree( + page_size=page_size, sliding_window_size=window + ) + kv = _swa_alloc(allocator, seq_len) + pool.write((0, slice(0, seq_len)), kv) + req = _make_req(0, list(range(seq_len)), 0, tree) + batch = _make_batch(tree, allocator, pool) + kwargs = {} if floor == "absent" else {"retain_floor": None} + free_swa_out_of_window_slots( + req, + seq_len - 1, + sliding_window_size=window, + page_size=page_size, + req_to_token_pool=batch.req_to_token_pool, + token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator, + **kwargs, + ) + frontiers.append(req.kv.swa_evicted_seqlen) + + expected = (seq_len - 1 - max(window, page_size)) // page_size * page_size + self.assertEqual(frontiers[0], expected) + self.assertEqual(frontiers[1], expected) + + def test_retain_floor_above_threshold_is_inert(self): + """The floor is a min(), so a checkpoint that is already inside the window + must not hold anything extra.""" + page_size, window = 8, 16 + seq_len = 200 + tree, allocator, pool = _build_swa_tree( + page_size=page_size, sliding_window_size=window + ) + kv = _swa_alloc(allocator, seq_len) + pool.write((0, slice(0, seq_len)), kv) + req = _make_req(0, list(range(seq_len)), 0, tree) + batch = _make_batch(tree, allocator, pool) + + free_swa_out_of_window_slots( + req, + seq_len - 1, + sliding_window_size=window, + page_size=page_size, + req_to_token_pool=batch.req_to_token_pool, + token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator, + retain_floor=seq_len, + ) + + expected = (seq_len - 1 - max(window, page_size)) // page_size * page_size + self.assertEqual(req.kv.swa_evicted_seqlen, expected) + + def test_retain_floor_does_not_unfree(self): + """The frontier only advances. A floor arriving after slots were already + freed must not claim them back, which would double-free on the next pass.""" + page_size, window = 8, 16 + tree, allocator, pool = _build_swa_tree( + page_size=page_size, sliding_window_size=window + ) + seq_len = 200 + kv = _swa_alloc(allocator, seq_len) + pool.write((0, slice(0, seq_len)), kv) + req = _make_req(0, list(range(seq_len)), 0, tree) + batch = _make_batch(tree, allocator, pool) + common_kwargs = dict( + sliding_window_size=window, + page_size=page_size, + req_to_token_pool=batch.req_to_token_pool, + token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator, + ) + + free_swa_out_of_window_slots(req, seq_len - 1, **common_kwargs) + advanced = req.kv.swa_evicted_seqlen + self.assertGreater(advanced, 0) + + free_swa_out_of_window_slots(req, seq_len - 1, retain_floor=0, **common_kwargs) + self.assertEqual(req.kv.swa_evicted_seqlen, advanced) + # -- Eviction formula: page_size == 1 -- def test_formula_page_size_1(self):