[mem_cache] Make free_swa sync-free on page_size == 1 (#36723)

This commit is contained in:
Liangsheng Yin
2026-09-02 14:18:22 -07:00
committed by GitHub
parent acea43079f
commit 19c7679e9e
4 changed files with 129 additions and 78 deletions
@@ -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."""
@@ -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: