diff --git a/python/sglang/srt/mem_cache/allocator/base.py b/python/sglang/srt/mem_cache/allocator/base.py index c9826e068..7001324dd 100644 --- a/python/sglang/srt/mem_cache/allocator/base.py +++ b/python/sglang/srt/mem_cache/allocator/base.py @@ -192,6 +192,23 @@ class BaseTokenToKVPoolAllocator(abc.ABC): Starts sit on page boundaries, ends may fall mid-page, and the page ranges of consecutive segments do not overlap -- so in page units the segments are aligned and disjoint, and every page is released once.""" + for free_index, start_pos in self._page_disjoint(segments): + self.free_segment(free_index, start_pos=start_pos) + + def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int): + """free_full() for a kv-row segment; same start-alignment contract as + free_segment(). Default: plain free_full().""" + assert start_pos % self.page_size == 0, ( + f"segment start {start_pos} is not page-aligned" + ) + self.free_full(free_index) + + def free_full_segments(self, segments): + """free_segments() for the full side alone; see free_full().""" + for free_index, start_pos in self._page_disjoint(segments): + self.free_full_segment(free_index, start_pos=start_pos) + + def _page_disjoint(self, segments): ps = self.page_size prev_end = None for free_index, start_pos in segments: @@ -202,4 +219,4 @@ class BaseTokenToKVPoolAllocator(abc.ABC): f"segment at {start_pos} shares a page with the one ending at {prev_end}" ) prev_end = start_pos + n - self.free_segment(free_index, start_pos=start_pos) + yield free_index, start_pos diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index bc4bd4262..8e27f9f56 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -103,7 +103,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self.release_pages = None self.free_group = None self.swa_free_group = [] - self.full_free_group = [] self._kvcache = kvcache self.clear() @@ -402,18 +401,34 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self.full_to_swa_index_mapping[free_index] == 0, msg="caller wants free", ) - if self.free_group is None: - self.full_attn_allocator.free(free_index) - else: - self.full_free_group.append(self._copy_for_free_group(free_index)) + self.full_attn_allocator.free(free_index) assert ( self.full_attn_allocator.available_size() <= self.full_attn_allocator.size ) + def free_segment(self, free_index: torch.Tensor, *, start_pos: int): + if free_index.numel() == 0: + return + # SWA first, as in free(): it reads the mapping that a later cache + # action in this group may re-point. + self.free_swa(free_index) + self.full_attn_allocator.free_segment(free_index, start_pos=start_pos) + + def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int): + if free_index.numel() == 0: + return + expect( + _SWA_PEER_RELEASED, + self.full_to_swa_index_mapping[free_index] == 0, + msg="caller wants free_segment", + ) + self.full_attn_allocator.free_segment(free_index, start_pos=start_pos) + def free_group_begin(self): super().free_group_begin() self.swa_free_group = [] - self.full_free_group = [] + # No full-side pile here: the full allocator's own group defers those. + self.full_attn_allocator.free_group_begin() def free_group_end(self): super().free_group_end() @@ -421,10 +436,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): swa_free_group = self.swa_free_group self.swa_free_group = [] self._release_swa(torch.cat(swa_free_group)) - if self.full_free_group: - full_free_group = self.full_free_group - self.full_free_group = [] - self.full_attn_allocator.free(torch.cat(full_free_group)) + self.full_attn_allocator.free_group_end() assert ( self.full_attn_allocator.available_size() <= self.full_attn_allocator.size ) @@ -468,7 +480,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): self.full_to_swa_index_mapping[:-1].fill_(0) self.free_group = None self.swa_free_group = [] - self.full_free_group = [] def get_cpu_copy(self, indices, mamba_indices=None): return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices) @@ -594,8 +605,16 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): # release once the SWA side is gone. return - # Not inherited: the SWA parent's hooks drive swa_free_group and - # full_free_group, which this pure-SWA variant does not have. + def free_segment(self, free_index: torch.Tensor, *, start_pos: int): + # Single pool: the parent's split into an SWA and a full half would + # release the same slots twice. + self.free(free_index) + + def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int): + return + + # Not inherited: the SWA parent's hooks drive swa_free_group and the full + # allocator's group, which this pure-SWA variant does not have. def free_group_begin(self): BaseTokenToKVPoolAllocator.free_group_begin(self) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 43cd1b15f..0d45290d6 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -127,7 +127,7 @@ def free_kv_row_segments( ) -> None: """Free ascending disjoint ``(kv_indices, start_pos)`` segments of one request's kv row, split at the SWA eviction floor.""" - swa_dead: list[torch.Tensor] = [] + swa_dead: list[tuple[torch.Tensor, int]] = [] swa_alive: list[tuple[torch.Tensor, int]] = [] for kv_indices, start_pos in segments: num_indices = kv_indices.numel() @@ -137,23 +137,19 @@ def free_kv_row_segments( # the deliberately unmapped prefix of a PD decode SWA-tail prealloc. num_dead = min(max(swa_evicted_seqlen - start_pos, 0), num_indices) if num_dead > 0: - swa_dead.append(kv_indices[:num_dead]) + swa_dead.append((kv_indices[:num_dead], start_pos)) if num_dead < num_indices: swa_alive.append((kv_indices[num_dead:], start_pos + num_dead)) if swa_dead and swa_alive: - # A mid-page floor would send a page shared by the dead and alive - # sides back twice. + # The two sides are separate calls, so neither one's page-disjointness + # check sees a floor that splits a page between them. assert swa_evicted_seqlen % allocator.page_size == 0, ( f"SWA eviction floor {swa_evicted_seqlen} splits a page " f"(page_size {allocator.page_size})" ) - if len(swa_dead) == 1: - allocator.free_full(swa_dead[0]) - elif swa_dead: - # Two dead pieces can share a boundary page, and only free_full's own - # page dedup covers that -- free_segments trims the alive side alone. - allocator.free_full(torch.cat(swa_dead)) + if swa_dead: + allocator.free_full_segments(swa_dead) if swa_alive: allocator.free_segments(swa_alive) diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index c6964dcab..5d6439366 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -3227,6 +3227,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): self.free_group = None self.free_page_reps_group: Optional[List[torch.Tensor]] = None + self.full_free_group: List[torch.Tensor] = [] # Empty (not None) for the leak checker. self.free_pages = torch.empty(0, dtype=torch.int64, device=device) self.release_pages = torch.empty(0, dtype=torch.int64, device=device) @@ -3661,6 +3662,17 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): self.full_attn_allocator.free(free_index.detach().to(torch.int64)) self.full_attn_allocator.clear_inverse_history() + def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None: + if free_index is None or free_index.numel() == 0: + return + if self.page_size == 1: + # token == page: free_full already frees by exact ids, no dedup. + self.free_full(free_index) + return + # The swa v2p is the mapping, so a tombstoned swa page drops out of the + # two-sided segment path by itself; full-only is the same call. + self.free_segment(free_index, start_pos=start_pos) + def set_full_to_swa_mapping( self, full_indices: torch.Tensor, swa_indices: torch.Tensor ) -> None: @@ -3676,13 +3688,20 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): # -- free-group -- + # Not the SWA parent's hooks: those open the parent's paged full allocator + # as a free group, and this composite's sub-pools defer on their own. def free_group_begin(self) -> None: - super().free_group_begin() + BaseTokenToKVPoolAllocator.free_group_begin(self) self.free_page_reps_group = [] + self.full_free_group = [] def free_group_end(self) -> None: pending, self.free_page_reps_group = self.free_page_reps_group, None - super().free_group_end() + full_free_group, self.full_free_group = self.full_free_group, [] + BaseTokenToKVPoolAllocator.free_group_end(self) + if full_free_group: + self.full_attn_allocator.free(torch.cat(full_free_group)) + self.full_attn_allocator.clear_inverse_history() if pending: self._release_page_reps(pending) @@ -3746,6 +3765,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): self.swa_attn_allocator.clear() self.free_group = None self.free_page_reps_group = None + self.full_free_group = [] # -- Lazy compaction hooks -- diff --git a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py index e76108615..e46a4cb05 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py @@ -503,10 +503,11 @@ class FullComponent(TreeComponent): if isinstance(action, FreeComponentDeviceSlot): alloc = self.cache.token_to_kv_pool_allocator for indices in action.indices: + # tree values are page-aligned copies of a kv row: page-exact segments if self.cache.is_swa_enabled: - alloc.full_attn_allocator.free(indices) + alloc.full_attn_allocator.free_segment(indices, start_pos=0) else: - alloc.free(indices) + alloc.free_segment(indices, start_pos=0) return raise AssertionError( f"FullComponent: unhandled ComponentAction {type(action).__name__}" 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 fd03d1b35..597e90cc0 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 @@ -1371,7 +1371,7 @@ class SWAComponent(TreeComponent): swa_value = self._translate_full_to_swa(action.incoming_full) alloc.set_full_to_swa_mapping(action.kept_full, swa_value) alloc.clear_full_to_swa_mapping(action.incoming_full) - alloc.free_full(action.incoming_full) + alloc.free_full_segment(action.incoming_full, start_pos=0) self.tree_core.set_component_device_value( action.node_id, self.component_type, swa_value ) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 9f10f0f31..38b236af7 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -1097,7 +1097,7 @@ class UnifiedRadixCache(BasePrefixCache): self.token_to_kv_pool_allocator.free_segment(indices, start_pos=0) elif isinstance(action, FreeDeviceKVFullOnly): for indices in action.indices: - self.token_to_kv_pool_allocator.free_full(indices) + self.token_to_kv_pool_allocator.free_full_segment(indices, start_pos=0) elif isinstance(action, BackupKV): if self.linker is not None: self.linker.offload_nodes(action.node_ids) diff --git a/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py index 1d07d1b99..9484cada8 100644 --- a/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py +++ b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py @@ -141,7 +141,7 @@ class _StaticAllocRecorder: self.clear_calls.append(full) self.full_to_swa_index_mapping[full.to(torch.int64)] = 0 - def free_full(self, indices): + def free_full_segment(self, indices, *, start_pos): self.freed_full.append(indices) def free(self, indices): @@ -307,8 +307,8 @@ class TestStaticPoolPathUnchanged(unittest.TestCase): ), "incoming ids' mapping entries must be zeroed (static recipe)", ) - # Through free_full, not the inner allocator: the latter skips the - # free-group defer. + # Through free_full_segment, not the inner allocator: the latter skips + # the free-group defer. self.assertEqual(len(static.freed_full), 1) self.assertEqual(static.freed_via_inner, []) ((node_id, ct, _),) = probe.tree_core.set_calls diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 0df9e7f74..10cf0487b 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -2,6 +2,7 @@ import unittest from array import array from types import SimpleNamespace from unittest import mock +from unittest.mock import patch import torch @@ -830,13 +831,13 @@ class TestSWA(unittest.TestCase): req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device) freed_lens = [] - original_free = allocator.free + original_free_segment = allocator.free_segment - def wrapped_free(indices): + def wrapped_free_segment(indices, *, start_pos): freed_lens.append(int(indices.numel())) - return original_free(indices) + return original_free_segment(indices, start_pos=start_pos) - allocator.free = wrapped_free + allocator.free_segment = wrapped_free_segment tree.cache_finished_req( req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len ) @@ -1070,21 +1071,45 @@ class TestFreeKvRow(CustomTestCase): ) self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline)) - def test_adjacent_below_floor_pieces_release_their_shared_page_once(self): + def test_below_floor_pieces_go_back_through_the_full_side(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. - free_kv_row_segments( - allocator, - [(indices[:6], 0), (indices[6:], 6)], - swa_evicted_seqlen=8, - ) + # Both rows [0, 4) and [4, 8) sit below the floor: full side only. + with patch.object( + allocator.full_attn_allocator, + "free", + side_effect=AssertionError("full side took the unique path"), + ): + free_kv_row_segments( + allocator, + [(indices[:4], 0), (indices[4:], 4)], + swa_evicted_seqlen=8, + ) self.assertEqual(allocator.full_available_size(), after_alloc + 8) + def test_grouped_full_side_frees_defer_and_skip_the_unique_path(self): + _, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) + indices = _swa_alloc(allocator, 12) + allocator.free_swa(indices[:8]) + after_alloc = allocator.full_available_size() + + with patch.object( + allocator.full_attn_allocator, + "free", + side_effect=AssertionError("full side took the unique path"), + ): + allocator.free_group_begin() + # dead rows [0, 8) and the alive row [8, 12) from one request + free_kv_row_segments(allocator, [(indices, 0)], swa_evicted_seqlen=8) + self.assertEqual(allocator.full_available_size(), after_alloc) + allocator.free_group_end() + + self.assertEqual(allocator.full_available_size(), after_alloc + 12) + def test_free_kv_row_reads_the_record_row_and_its_floor(self): indices = _swa_alloc(self.allocator, 8) cache = _RowCache(self.allocator, indices) diff --git a/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py index 364fdde98..da61f4a83 100644 --- a/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py +++ b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py @@ -406,6 +406,26 @@ class TestEveryUnifiedAllocatorOverridesFreeSegment(unittest.TestCase): self.assertIn("free_page_reps_group", inspect.getsource(cls)) +class TestUnifiedSwaFullSideGroup(unittest.TestCase): + """`free_full` inside a free group must defer and land at `free_group_end`; + the composite carries that pile itself, not through the SWA parent's hooks.""" + + def test_full_only_frees_defer_until_group_end(self): + from test_unified_swa_shared_virtual_ids import _build + + alloc = _build(64, 32, 2, 4) + idx = alloc.alloc(8) + alloc.free_swa(idx) + before = alloc.full_available_size() + + alloc.free_group_begin() + alloc.free_full(idx[:4]) + alloc.free_full_segment(idx[4:], start_pos=4) + self.assertEqual(alloc.full_available_size(), before) + alloc.free_group_end() + self.assertEqual(alloc.full_available_size(), before + 8) + + class TestFreeSwaWindowRatchetNoHostSync(unittest.TestCase): """The per-decode-step SWA window ratchet frees a CONTIGUOUS row slice with host-int, page-aligned bounds — the same shape `free_segment` was @@ -471,6 +491,25 @@ class TestFreeSwaWindowRatchetNoHostSync(unittest.TestCase): alloc.free_swa(v[: 4 * self.PS], start_pos=0) alloc.free_swa(v[4 * self.PS :], start_pos=4 * self.PS) + def test_full_only_segment_free_never_syncs(self): + """The request-finish dead half (swa already tombstoned) frees the full + side by page reps: no unique from free_full's token dedup.""" + alloc = self._swa_composite(lazy=True) + v = alloc.alloc(8 * self.PS) + alloc.free_swa(v, start_pos=0) + before = alloc.full_available_size() + with ( + mock.patch.object( + torch, "unique", side_effect=AssertionError("unique = host sync") + ), + mock.patch.object( + torch.Tensor, "item", side_effect=AssertionError("item = host sync") + ), + ): + alloc.free_full_segment(v[: 4 * self.PS], start_pos=0) + alloc.free_full_segment(v[4 * self.PS :], start_pos=4 * self.PS) + self.assertEqual(alloc.full_available_size(), before + 8 * self.PS) + def test_unaligned_start_pos_is_rejected(self): """A mid-page start must fail loudly, not release the head page whole.""" alloc = self._swa_composite(lazy=True) 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 5c55c7684..5bc947289 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 @@ -516,15 +516,18 @@ def build_fixture( full_attention_layer_ids=cfg.full_attention_layer_ids, device=device, ) - allocator = SWATokenToKVPoolAllocator( - size=cfg.kv_size, - size_swa=cfg.kv_size, - page_size=cfg.page_size, - dtype=cfg.dtype, - device=device, - kvcache=kv_pool, - need_sort=False, - ) + # Tree values reach the allocator as start_pos=0 segments, so only the + # debug cross-check against torch.unique catches a mis-aligned value. + with envs.SGLANG_DEBUG_MEMORY_POOL.override(True): + allocator = SWATokenToKVPoolAllocator( + size=cfg.kv_size, + size_swa=cfg.kv_size, + page_size=cfg.page_size, + dtype=cfg.dtype, + device=device, + kvcache=kv_pool, + need_sort=False, + ) elif cfg.has_mamba: kv_pool = HybridLinearKVPool( size=cfg.kv_size, @@ -4070,11 +4073,8 @@ class UnifiedRadixCacheSuite: # for an un-charged gate to accept. extend_need = 2 * ps + 1 max_new = 8 - self.assertIsNotNone( - cons_alloc.swa_attn_allocator.alloc( - cons_alloc.swa_available_size() - (window + extend_need - 1) - ) - ) + hold = cons_alloc.swa_available_size() - (window + extend_need - 1) + self.assertIsNotNone(cons_alloc.swa_attn_allocator.alloc(hold // ps * ps)) # The adder's SWA gate for this request (_swa_budget_for_req). surfaced_swa_hit = cons.staged_prefetch_swa_tokens(req_id) @@ -7559,8 +7559,8 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase): _component_with_cache(ComponentType.FULL, cache).apply_component_action( FreeComponentDeviceSlot([indices], component_type=ComponentType.FULL) ) - cache.token_to_kv_pool_allocator.full_attn_allocator.free.assert_called_once_with( - indices + cache.token_to_kv_pool_allocator.full_attn_allocator.free_segment.assert_called_once_with( + indices, start_pos=0 ) def test_apply_component_action_device_kv_swa_uses_free_swa(self): @@ -7681,7 +7681,7 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase): alloc.set_full_to_swa_mapping.assert_called_once_with(kept_full, swa_value) # the incoming full's stale mapping is cleared, then its slot freed (full-only) alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full) - alloc.free_full.assert_called_once_with(incoming_full) + alloc.free_full_segment.assert_called_once_with(incoming_full, start_pos=0) # Never by indexing the tensor: the unified composite has no # `full_to_swa_index_mapping` to index into. alloc.full_to_swa_index_mapping.__setitem__.assert_not_called()