[mem_cache] Make free_swa sync-free on page_size == 1 (#36723)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user