diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index b7f94ea55..bc4bd4262 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -6,6 +6,7 @@ from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.utils import is_npu from sglang.srt.utils.common import get_num_new_pages +from sglang.srt.utils.invariants import Bucket, Invariant, IsTrue, expect _is_npu = is_npu() @@ -17,6 +18,13 @@ if _is_npu: ) +# free_swa releases whatever the mapping points at, so an entry that reads as the +# padding slot would push slot 0 into the SWA free list and hand it out twice. +_SWA_PEER_MAPPED = Invariant("swa.peer_mapped", Bucket.FATAL_UNCONTAINABLE, IsTrue()) +# free_full leaves the mapping alone, so a live entry would strand its SWA peer. +_SWA_PEER_RELEASED = Invariant("swa.peer_released", Bucket.GUARD, IsTrue()) + + class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): """Allocator for SWA hybrid KV cache.""" @@ -355,11 +363,15 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): return if self.page_size == 1: + # A filter here would make the output shape data-dependent, + # which costs a device-to-host sync. mapping_indices = free_index + swa_indices = self.full_to_swa_index_mapping[mapping_indices] + expect(_SWA_PEER_MAPPED, swa_indices > 0, msg="caller wants free_full") else: mapping_indices = self._expand_to_full_pages(free_index) + swa_indices = self.full_to_swa_index_mapping[mapping_indices] - swa_indices = self.full_to_swa_index_mapping[mapping_indices] self.clear_full_to_swa_mapping(mapping_indices) if self.free_group is not None: @@ -371,18 +383,26 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self._release_swa(swa_indices) def _release_swa(self, swa_indices: torch.Tensor): - # One filter per group: its data-dependent shape costs a sync, and - # filtering the batch selects the same slots as filtering per call. - self.swa_attn_allocator.free(swa_indices[swa_indices > 0]) + if self.page_size > 1: + # HiCache LOAD_BACK re-pairs a page-aligned full chunk with an offset + # SWA one (commit_hicache_transfer advances by raw token count), so a + # page can hold unmapped slots; one filter per group, not per call. + swa_indices = swa_indices[swa_indices > 0] + self.swa_attn_allocator.free(swa_indices) assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size def free_full(self, free_index: torch.Tensor): if free_index.numel() == 0: return + # Checked at enqueue: a cache action later in this group may pair the + # slot again, and that new peer is not this call's to judge. + expect( + _SWA_PEER_RELEASED, + self.full_to_swa_index_mapping[free_index] == 0, + msg="caller wants free", + ) if self.free_group is None: - # Full side only: a tombstoned range's mapping entries read as the - # padding slot, so `free` would push slot 0 into the SWA free list. self.full_attn_allocator.free(free_index) else: self.full_free_group.append(self._copy_for_free_group(free_index)) @@ -404,7 +424,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): if self.full_free_group: full_free_group = self.full_free_group self.full_free_group = [] - self.free_full(torch.cat(full_free_group)) + self.full_attn_allocator.free(torch.cat(full_free_group)) assert ( self.full_attn_allocator.available_size() <= self.full_attn_allocator.size ) diff --git a/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py b/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py index ecfc24e04..42509ba95 100644 --- a/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py +++ b/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py @@ -22,6 +22,9 @@ class _PerturbReqToTokenBase(CanaryE2EBase): # still looks busy). That's expected for this test; disable strict # mode so the leak warning doesn't crash the scheduler. "SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "0", + # A perturbed slot reaches free_swa with its peer already released: that + # is the corruption under test, not a swa.peer_mapped regression. + "SGLANG_INVARIANT_CHECK": "0", } @classmethod diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 3503faf4d..cf6267b1d 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -1,11 +1,12 @@ import unittest from array import array from types import SimpleNamespace +from unittest import mock import torch from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored -from sglang.srt.environ import envs +from sglang.srt.environ import InvariantCheckLevel, envs from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.swa import ( PureSWATokenToKVPoolAllocator, @@ -109,6 +110,20 @@ def _build_swa_tree( return tree, allocator, req_to_token_pool +def _sync_error(fn): + """The RuntimeError torch raises if `fn` synchronizes, or None.""" + torch.cuda.synchronize() + torch.cuda.set_sync_debug_mode("error") + try: + fn() + except RuntimeError as exc: + return exc + finally: + torch.cuda.set_sync_debug_mode("default") + torch.cuda.synchronize() + return None + + def _build_pure_swa_allocator(size_swa: int = 16): device = get_device() kv_pool = SWAKVPool( @@ -270,28 +285,16 @@ class TestSWA(unittest.TestCase): # its own, which the detector would report as this call's fault. allocator.clear_full_to_swa_mapping(full_indices) - def sync_error(fn): - torch.cuda.synchronize() - torch.cuda.set_sync_debug_mode("error") - try: - fn() - except RuntimeError as exc: - return exc - finally: - torch.cuda.set_sync_debug_mode("default") - torch.cuda.synchronize() - return None - # Gate on the pre-fix form: a detector blind to this sync class would pass # the assert below no matter how the mapping is cleared. - pre_fix_error = sync_error( + pre_fix_error = _sync_error( lambda: mapping.__setitem__(full_indices.to(torch.int64), 0) ) if pre_fix_error is None: self.skipTest("sync debug mode does not flag a blocking H2D copy here") self.assertIsNone( - sync_error(lambda: allocator.clear_full_to_swa_mapping(full_indices)) + _sync_error(lambda: allocator.clear_full_to_swa_mapping(full_indices)) ) def test_free_swa_group_owns_deferred_indices(self): @@ -422,19 +425,6 @@ class TestSWA(unittest.TestCase): ) self.assertEqual(num_mapped, num_in_use) - def test_free_group_owns_tombstoned_indices(self): - """free_swa then free of the same full slot must free the SWA slot once.""" - allocator, full_indices, _, swa_indices, _ = self._build_two_mapped_slots() - swa_available_before = allocator.swa_available_size() - - allocator.free_group_begin() - allocator.free_swa(full_indices) - allocator.free(full_indices) - allocator.free_group_end() - - self.assertEqual(allocator.swa_available_size(), swa_available_before + 1) - self.assertTrue(self._swa_slot_is_free(allocator, swa_indices)) - def test_pure_swa_rejects_mapping_edits(self): allocator = _build_pure_swa_allocator() indices = allocator.alloc(2) @@ -1006,26 +996,18 @@ class TestFreeFullPartition(CustomTestCase): self.allocator.swa_available_size(), ) - def test_free_full_keeps_the_swa_peers_allocated(self): + def test_free_full_touches_only_the_full_pool(self): indices = _swa_alloc(self.allocator, 4) + # free_full's precondition: the SWA peers are already released. + self.allocator.free_swa(indices) + self.assertEqual(self._sizes(), (self.full_baseline - 4, self.swa_baseline)) + self.allocator.free_full(indices) - - full_avail, swa_avail = self._sizes() - self.assertEqual(full_avail, self.full_baseline) - self.assertEqual(swa_avail, self.swa_baseline - 4) - - def test_free_full_leaves_the_mapping_intact(self): - indices = _swa_alloc(self.allocator, 4) - before = self.allocator.full_to_swa_index_mapping[indices].clone() - self.allocator.free_full(indices) - - self.assertTrue(bool((before > 0).all())) - self.assertTrue( - torch.equal(self.allocator.full_to_swa_index_mapping[indices], before) - ) + self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline)) def test_free_full_is_deferred_inside_a_free_group(self): indices = _swa_alloc(self.allocator, 4) + self.allocator.free_swa(indices) self.allocator.free_group_begin() self.allocator.free_full(indices) @@ -1062,7 +1044,7 @@ class TestFreeKvRow(CustomTestCase): self.allocator.swa_available_size(), ) - def test_floor_decides_how_much_of_the_swa_side_stays_out(self): + def test_floor_decides_how_much_of_the_swa_side_the_row_frees(self): # (start_pos, num_slots, floor, rows whose SWA peers are already gone) cases = [ (0, 4, 4, 4), @@ -1073,21 +1055,25 @@ class TestFreeKvRow(CustomTestCase): for start_pos, num_slots, floor, num_dead in cases: with self.subTest(start_pos=start_pos, floor=floor): indices = _swa_alloc(self.allocator, num_slots) + # Window eviction already released the peers below the floor. + if num_dead: + self.allocator.free_swa(indices[:num_dead]) + self.assertEqual( + self._sizes(), + ( + self.full_baseline - num_slots, + self.swa_baseline - num_slots + num_dead, + ), + ) free_kv_row_segments( self.allocator, [(indices, start_pos)], swa_evicted_seqlen=floor ) - self.assertEqual( - self._sizes(), - (self.full_baseline, self.swa_baseline - num_dead), - ) - # Give the held-back SWA peers back, so the next case starts clean. - if num_dead: - self.allocator.free_swa(indices[:num_dead]) self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline)) def test_adjacent_below_floor_pieces_release_their_shared_page_once(self): _, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) indices = _swa_alloc(allocator, 8) + allocator.free_swa(indices) after_alloc = allocator.full_available_size() # Rows [0, 6) and [6, 8) both sit below the floor and share page 1. @@ -1103,12 +1089,13 @@ class TestFreeKvRow(CustomTestCase): indices = _swa_alloc(self.allocator, 8) cache = _RowCache(self.allocator, indices) kv = SimpleNamespace(req_pool_idx=0, swa_evicted_seqlen=3) + self.allocator.free_swa(indices[:3]) cache.free_kv_row(kv, [(1, 5)]) - # Rows [1, 5) go back on the full side; of those, [1, 3) lost their SWA - # peers already, so 6 of the 8 SWA slots are still out. - self.assertEqual(self._sizes(), (self.full_baseline - 4, self.swa_baseline - 6)) + # Rows [1, 5) go back on the full side; only [3, 5) still had SWA peers + # to give back, so rows 5-7 keep the 3 SWA slots that are still out. + self.assertEqual(self._sizes(), (self.full_baseline - 4, self.swa_baseline - 3)) def test_single_pool_free_kv_row_still_frees_the_whole_range(self): allocator = _SinglePoolAllocator() @@ -1125,6 +1112,52 @@ class TestFreeKvRow(CustomTestCase): self.assertEqual(len(allocator.freed), 2) +class TestSWAPeerMappedContract(CustomTestCase): + """page_size 1 gives back every peer the mapping names, without filtering: + the contract replaces what `swa_indices > 0` used to absorb.""" + + def _strict(self): + return envs.SGLANG_INVARIANT_CHECK.override(int(InvariantCheckLevel.STRICT)) + + def _condition_checked_by(self, allocator, indices): + """The predicate free_swa hands the async assert, as a python bool.""" + with self._strict(): + with mock.patch.object(torch, "_assert_async") as assert_async: + allocator.free_swa(indices) + return bool(assert_async.call_args.args[0]) + + def test_free_swa_flags_a_slot_whose_peer_is_already_gone(self): + _, allocator, _ = _build_swa_tree(is_eagle=False) + live = _swa_alloc(allocator, 4) + stale = _swa_alloc(allocator, 4) + # Whoever released the peer left the mapping reading as the padding slot. + allocator.clear_full_to_swa_mapping(stale) + + self.assertTrue(self._condition_checked_by(allocator, live)) + self.assertFalse(self._condition_checked_by(allocator, stale)) + + @unittest.skipUnless(torch.cuda.is_available(), "sync detection needs CUDA") + def test_free_swa_does_not_synchronize(self): + """The filter's output shape was data-dependent, so it read a count back + to the host; the gather that replaced it has a fixed shape.""" + _, allocator, _ = _build_swa_tree(is_eagle=False) + mapping = allocator.full_to_swa_index_mapping + + # Warm up outside the window: a first-time cudaMalloc can synchronize on + # its own, which the detector would report as this call's fault. + allocator.free_swa(_swa_alloc(allocator, 4)) + indices = _swa_alloc(allocator, 4) + + # Gate on the pre-fix form: a detector blind to this sync class would pass + # the assert below no matter how free_swa reads the mapping. + peers = mapping[indices] + if _sync_error(lambda: peers[peers > 0]) is None: + self.skipTest("sync debug mode does not flag a data-dependent shape here") + + with self._strict(): + self.assertIsNone(_sync_error(lambda: allocator.free_swa(indices))) + + class TestCacheUnfinishedReqEvictedPrefix(CustomTestCase): """An unfinished request whose SWA prefix is already gone must insert that prefix as a tombstone, not as live SWA KV.""" 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 b00c05e2e..a65bfe49c 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 @@ -8007,13 +8007,10 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite): cache, allocator, _ = build_fixture(self.cfg) seq = list(range(1, self.cfg.sliding_window_size + 1)) key = RadixKey(array("q", seq)) - cache.insert( - InsertParams( - key=key, - value=self._alloc(allocator, len(seq)), - swa_evicted_seqlen=len(seq), - ) - ) + evicted = self._alloc(allocator, len(seq)) + # Window eviction already released the peers below the floor. + allocator.free_swa(evicted) + cache.insert(InsertParams(key=key, value=evicted, swa_evicted_seqlen=len(seq))) (leaf,) = _node_children(cache, cache.root_node_handle()) lock_result = cache.inc_lock_ref(leaf) if lock_full else None try: @@ -8045,11 +8042,9 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite): cache, allocator, _ = build_fixture(self.cfg) seq = list(range(1, 2 * sw + 1)) key = RadixKey(array("q", seq)) - cache.insert( - InsertParams( - key=key, value=self._alloc(allocator, len(seq)), swa_evicted_seqlen=sw - ) - ) + evicted = self._alloc(allocator, len(seq)) + allocator.free_swa(evicted[:sw]) + cache.insert(InsertParams(key=key, value=evicted, swa_evicted_seqlen=sw)) value = self._alloc(allocator, len(seq)) full_available = allocator.full_attn_allocator.available_size() swa_available = allocator.swa_attn_allocator.available_size() @@ -8073,11 +8068,9 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite): cache, allocator, req_to_token_pool = build_fixture(self.cfg) seq = list(range(1, 2 * sw + 1)) key = RadixKey(array("q", seq)) - cache.insert( - InsertParams( - key=key, value=self._alloc(allocator, len(seq)), swa_evicted_seqlen=sw - ) - ) + evicted = self._alloc(allocator, len(seq)) + allocator.free_swa(evicted[:sw]) + cache.insert(InsertParams(key=key, value=evicted, swa_evicted_seqlen=sw)) (prefix_node,) = _node_children(cache, cache.root_node_handle()) (window_node,) = _node_children(cache, prefix_node) self.assertIsNone(_device_value(cache, prefix_node, ComponentType.SWA)) @@ -8107,6 +8100,8 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite): cache.insert(InsertParams(key=key, value=self._alloc(allocator, len(seq)))) value = self._alloc(allocator, len(seq)) + # Window eviction already released the peers below the floor. + allocator.free_swa(value[:sw]) with mock.patch.object( cache, "_apply_cache_action", wraps=cache._apply_cache_action ) as spy: