Retain SWA down to the last state checkpoint (#34729)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user