diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index 676f2450a..29afd8822 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -431,8 +431,9 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): def free_swa(self, free_index: torch.Tensor): """Release the SWA peers of an arbitrary slot set and clear their mapping. - Synchronizes at page_size > 1; kv-row segments go through free_swa_segment().""" - if free_index.numel() == 0: + No-op for a per-request ring, which owns no paged SWA peers. Otherwise + synchronizes at page_size > 1; kv-row segments use free_swa_segment().""" + if self._swa_req_ring or free_index.numel() == 0: return if self.page_size == 1: @@ -455,8 +456,9 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): def free_swa_segment(self, free_index: torch.Tensor, *, start_pos: int): """free_swa() for a kv-row segment; same start-alignment contract as - free_segment(), and fixed-shape at every page size.""" - if free_index.numel() == 0: + free_segment(), and fixed-shape at every page size. No-op for a + per-request ring, as in free_swa().""" + if self._swa_req_ring or free_index.numel() == 0: return self._free_swa_pages(free_index, start_pos=start_pos) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py index bf8d90fcb..5ccb585b7 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py @@ -239,34 +239,34 @@ def _build_deepseek_v4_device_pool_group( ) -> DevicePoolGroup: from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( + _dsv4_compressed_region_buffers, _dsv4_indexer_regions, _resolve_deepseek_v4_layer_mappings, ) - mappings = _resolve_deepseek_v4_layer_mappings(kvcache) - if getattr(kvcache, "_unified_kv", False) or isinstance( - kvcache.c4_kv_pool, HiSparseC4DevicePool - ): - raise ValueError( - "The direct external linker does not support unified-KV or HiSparse." - ) - if kvcache.swa_page_size != page_size: - raise ValueError( - "DeepSeek V4 SWA page size must match the tree page size: " - f"{kvcache.swa_page_size} != {page_size}." - ) + if isinstance(kvcache.c4_kv_pool, HiSparseC4DevicePool): + raise ValueError("The direct external linker does not support HiSparse.") - entries = [ - DevicePoolEntry( - name=PoolName.SWA, - indices_from_pool=PoolName.SWA, - device_pool=kvcache.swa_kv_pool, - components=[kvcache.swa_kv_pool.kv_buffer], - layer_mapping=mappings.swa, - page_size=page_size, - rows_are_pages=True, + mappings = _resolve_deepseek_v4_layer_mappings(kvcache) + is_unified_kv = getattr(kvcache, "_unified_kv", False) + entries = [] + if not is_unified_kv: + if kvcache.swa_page_size != page_size: + raise ValueError( + "DeepSeek V4 SWA page size must match the tree page size: " + f"{kvcache.swa_page_size} != {page_size}." + ) + entries.append( + DevicePoolEntry( + name=PoolName.SWA, + indices_from_pool=PoolName.SWA, + device_pool=kvcache.swa_kv_pool, + components=[kvcache.swa_kv_pool.kv_buffer], + layer_mapping=mappings.swa, + page_size=page_size, + rows_are_pages=True, + ) ) - ] def add(name, source, pool, buffers, layer_mapping): if layer_mapping: @@ -282,11 +282,14 @@ def _build_deepseek_v4_device_pool_group( ) ) + c4_buffers, _ = _dsv4_compressed_region_buffers(kvcache, 4) + c128_buffers, _ = _dsv4_compressed_region_buffers(kvcache, 128) + add( PoolName.DEEPSEEK_V4_C4, PoolName.KV, kvcache.c4_kv_pool, - kvcache.c4_kv_pool.kv_buffer, + c4_buffers, mappings.c4, ) for region in _dsv4_indexer_regions(kvcache, page_size): @@ -301,29 +304,30 @@ def _build_deepseek_v4_device_pool_group( PoolName.DEEPSEEK_V4_C128, PoolName.KV, kvcache.c128_kv_pool, - kvcache.c128_kv_pool.kv_buffer, + c128_buffers, mappings.c128, ) - add( - PoolName.DEEPSEEK_V4_C4_STATE, - PoolName.SWA, - kvcache.compress_state_pools, - _deepseek_v4_state_views( + if not is_unified_kv: + add( + PoolName.DEEPSEEK_V4_C4_STATE, + PoolName.SWA, kvcache.compress_state_pools, - mappings.c4_state_global_layers, - ), - mappings.c4_state, - ) - add( - PoolName.DEEPSEEK_V4_C4_INDEXER_STATE, - PoolName.SWA, - kvcache.indexer_compress_state_pools, - _deepseek_v4_state_views( + _deepseek_v4_state_views( + kvcache.compress_state_pools, + mappings.c4_state_global_layers, + ), + mappings.c4_state, + ) + add( + PoolName.DEEPSEEK_V4_C4_INDEXER_STATE, + PoolName.SWA, kvcache.indexer_compress_state_pools, - mappings.c4_state_global_layers, - ), - mappings.c4_state, - ) + _deepseek_v4_state_views( + kvcache.indexer_compress_state_pools, + mappings.c4_state_global_layers, + ), + mappings.c4_state, + ) return DevicePoolGroup( entries, mappings.transfer_layer_num, 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 3d38ee282..6878bc03f 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 @@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Callable, Optional, Sequence import torch from sglang.srt.environ import envs +from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring from sglang.srt.mem_cache.base_prefix_cache import ( DecLockRefParams, EvictParams, @@ -308,11 +309,9 @@ class SWAComponent(TreeComponent): ct = self.component_type state = {"len": float("inf")} - # unified_kv never caches the SWA ring (per-request, not content-stable), - # so SWA bookkeeping must not gate the match here. - swa_device_only_hicache = ( - not self.tree_core.has_swa_host_pool and self.tree_core.enable_hicache - ) + # A per-request SWA ring is not stored in tree nodes, so its bookkeeping + # must not gate prefix matching. + swa_req_ring = is_swa_req_ring(self.cache.token_to_kv_pool_allocator) def validator(node: UnifiedTreeNode) -> bool: cd = node.component_data[ct] @@ -320,7 +319,7 @@ class SWAComponent(TreeComponent): # — load_back will restore SWA from host before use. if cd.value is None and (match_device_only or cd.host_value is None): state["len"] = 0 - if swa_device_only_hicache and (node.backuped or not node.evicted): + if swa_req_ring and (node.backuped or not node.evicted): return True return False state["len"] += len(node.key) diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py b/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py index 5c163a676..8304df0eb 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py @@ -25,6 +25,7 @@ from typing import TYPE_CHECKING, NamedTuple import torch +from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring from sglang.srt.mem_cache.base_prefix_cache import ( DecLockRefParams, InsertParams, @@ -160,6 +161,15 @@ class UnifiedCacheLinkerWrapper: self.cache = cache self.cache_linker = cache_linker + swa = cache.components.get(ComponentType.SWA) + self._skip_swa = swa is not None and is_swa_req_ring( + cache.token_to_kv_pool_allocator + ) + self._components = tuple( + component + for component in cache._components_tuple + if not (self._skip_swa and component is swa) + ) # rid -> what match found, consumed by the next init_load_back. self.hit_markers: dict[str, ExternalCacheHitMarker] = {} # Loads in flight, each pinning its inserted endpoint until DMA completes. @@ -192,7 +202,7 @@ class UnifiedCacheLinkerWrapper: return result lookup_transfers = [] - for component in cache._components_tuple: + for component in self._components: transfer = component.build_external_linker_transfer( LinkerTransferPhase.LOOKUP, None, tail_hashes ) @@ -290,7 +300,7 @@ class UnifiedCacheLinkerWrapper: # Build per-component linker transfers. component_transfers: list[tuple[TreeComponent, PoolTransfer]] = [] - for component in cache._components_tuple: + for component in self._components: transfer = component.build_external_linker_transfer( LinkerTransferPhase.LOAD, None, tail_hashes ) @@ -313,6 +323,20 @@ class UnifiedCacheLinkerWrapper: prefix_len, ) + # Components omitted from the linker do not run their PREPARE hook. + # Keep a non-restorable SWA range as tombstones instead of rebuilding + # it from an uninitialized FULL-to-SWA mapping during cache.insert(). + if self._skip_swa: + if req.kv is None: + from sglang.srt.managers.schedule_batch import ReqKvInfo + + req.kv = ReqKvInfo( + kv_allocated_len=prefix_len, + swa_evicted_seqlen=prefix_len, + ) + else: + req.kv.swa_evicted_seqlen = max(req.kv.swa_evicted_seqlen, prefix_len) + # Insert the newly loaded tail into the tree. prefix_indices = torch.cat( [req.prefix_indices.to(torch.int64), full_transfer.device_indices] @@ -484,6 +508,8 @@ class UnifiedCacheLinkerWrapper: node_id ) if transfers is not None: + if self._skip_swa: + transfers = [t for t in transfers if t.name != PoolName.SWA] self._offload_node(node_id, transfers) def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None: diff --git a/test/registered/unit/mem_cache/test_linker_pool_assembler.py b/test/registered/unit/mem_cache/test_linker_pool_assembler.py index e26d597f2..00547a140 100644 --- a/test/registered/unit/mem_cache/test_linker_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_linker_pool_assembler.py @@ -2,6 +2,7 @@ import unittest from types import SimpleNamespace +from unittest.mock import Mock, call import torch @@ -13,6 +14,7 @@ from sglang.srt.mem_cache.hicache_storage import ( from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import ( DevicePoolEntry, DevicePoolGroup, + _build_deepseek_v4_device_pool_group, resolve_hybrid_device_pool_group, ) from sglang.srt.mem_cache.unified_cache.component_type import ComponentType @@ -180,7 +182,8 @@ class TestHybridDevicePoolAssembler(CustomTestCase): kv_buffer=[torch.zeros((8, 3), dtype=torch.uint8) for _ in range(3)] ) kvcache.c4_kv_pool = SimpleNamespace( - kv_buffer=[torch.zeros((8, 5), dtype=torch.uint8) for _ in range(2)] + kv_buffer=[torch.zeros((8, 5), dtype=torch.uint8) for _ in range(2)], + bytes_per_page_padded=5, ) kvcache.c4_indexer_kv_pool = SimpleNamespace( index_k_with_scale_buffer=[ @@ -188,7 +191,8 @@ class TestHybridDevicePoolAssembler(CustomTestCase): ] ) kvcache.c128_kv_pool = SimpleNamespace( - kv_buffer=[torch.zeros((8, 11), dtype=torch.uint8)] + kv_buffer=[torch.zeros((8, 11), dtype=torch.uint8)], + bytes_per_page_padded=11, ) kvcache.layer_mapping = [ DeepSeekV4LayerItem(0, -1), @@ -236,6 +240,95 @@ class TestHybridDevicePoolAssembler(CustomTestCase): self.assertEqual(offsets, [[5]]) self.assertIsNone(c4_pool.get_prepared_layer_range_meta([0], 1)) + def test_unified_deepseek_v4_uses_only_compressed_pools(self): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4LayerItem + + for split_indexer in (False, True): + with self.subTest(split_indexer=split_indexer): + c4 = [torch.zeros((4, width), dtype=torch.uint8) for width in (5, 7)] + c128 = [torch.zeros((4, 11), dtype=torch.uint8)] + expected = { + PoolName.DEEPSEEK_V4_C4: c4, + PoolName.DEEPSEEK_V4_C128: c128, + } + if split_indexer: + payload = [ + torch.zeros((4, 1, 4, 1, 16), dtype=torch.uint8) + for _ in range(2) + ] + scale = [ + torch.zeros((4, 1, 4, 1), dtype=torch.uint8) for _ in range(2) + ] + indexer = SimpleNamespace( + index_k_with_scale_buffer=None, + index_k_payload_buffer=payload, + index_k_scale_buffer=scale, + ) + expected[PoolName.DEEPSEEK_V4_C4_INDEXER] = [ + b.flatten(1) for b in payload + ] + expected[PoolName.DEEPSEEK_V4_C4_INDEXER_SCALE] = [ + b.flatten(1) for b in scale + ] + else: + buffers = [ + torch.zeros((4, width), dtype=torch.uint8) for width in (13, 17) + ] + indexer = SimpleNamespace(index_k_with_scale_buffer=buffers) + expected[PoolName.DEEPSEEK_V4_C4_INDEXER] = buffers + regions = {4: (c4, 7), 128: (c128, 11)} + kvcache = SimpleNamespace( + _unified_kv=True, + start_layer=0, + end_layer=3, + layer_mapping=[ + DeepSeekV4LayerItem(4, 1), + DeepSeekV4LayerItem(128, 0), + DeepSeekV4LayerItem(4, 0), + ], + # Unified KV has neither paged KV nor an index-addressed SWA pool. + swa_kv_pool=None, + c4_kv_pool=None, + c128_kv_pool=None, + swa_page_size=3, + c4_indexer_kv_pool=indexer, + unified_region_buffers=Mock(side_effect=regions.__getitem__), + ) + group = _build_deepseek_v4_device_pool_group(kvcache, page_size=2) + + self.assertEqual(set(group.entry_map), set(expected)) + self.assertEqual(set(group.sources.values()), {PoolName.KV}) + self.assertTrue(group.rank_replicated) + for name, buffers in expected.items(): + entry = group.entry_map[name] + actual = entry.components[0] + self.assertEqual(len(actual), len(buffers)) + for got, want in zip(actual, buffers): + self.assertEqual(got.data_ptr(), want.data_ptr()) + self.assertEqual(got.shape, want.shape) + self.assertEqual( + kvcache.unified_region_buffers.call_args_list, [call(4), call(128)] + ) + resolved = group.resolve_transfers( + [ + PoolTransfer( + name=PoolName.KV, + keys=["page-0"], + device_indices=torch.tensor([0, 1]), + ) + ] + ) + self.assertEqual({t.name for t in resolved}, set(expected)) + + def test_deepseek_v4_still_rejects_hisparse(self): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool + + kvcache = SimpleNamespace( + c4_kv_pool=HiSparseC4DevicePool.__new__(HiSparseC4DevicePool) + ) + with self.assertRaisesRegex(ValueError, "does not support HiSparse"): + _build_deepseek_v4_device_pool_group(kvcache, 2) + def test_dsa_uses_hybrid_assembler_strategy(self): from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 7b5cd7b63..65c9b815e 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -58,6 +58,7 @@ def _build_swa_tree( kv_size_swa: int = 32, sliding_window_size: int = 4, enable_kv_cache_events: bool = False, + swa_req_ring_size: int | None = None, ): head_num = 8 head_dim = 128 @@ -88,6 +89,7 @@ def _build_swa_tree( full_attention_layer_ids=full_attention_layer_ids, device=device, ) + kv_pool.swa_req_ring_size = swa_req_ring_size allocator = SWATokenToKVPoolAllocator( size=kv_size, size_swa=kv_size_swa, @@ -96,6 +98,7 @@ def _build_swa_tree( device=device, kvcache=kv_pool, need_sort=False, + req_to_token_pool=req_to_token_pool, ) tree = SWARadixCache( params=CacheInitParams( @@ -1234,6 +1237,125 @@ class TestSWAPeerMappedContract(CustomTestCase): self.assertIsNone(_sync_error(lambda: allocator.free_swa(indices))) +@unittest.skipUnless(torch.cuda.is_available(), "paged allocation kernels need CUDA") +class TestSWAReqRingFree(CustomTestCase): + PS = 256 + + def _allocated_ring(self): + ps = self.PS + _, allocator, req_pool = _build_swa_tree( + is_eagle=False, + page_size=ps, + req_size=2, + max_context_len=4 * ps, + kv_size=4 * ps, + kv_size_swa=2 * ps, + swa_req_ring_size=ps, + ) + self.assertTrue(allocator.swa_req_ring) + self.assertIsNotNone(req_pool.alloc_rows(1)) + device = allocator.device + prefix_cpu = torch.tensor([0], dtype=torch.int64) + seq_cpu = torch.tensor([2 * ps], dtype=torch.int64) + # Use the real ring allocation paths: only FULL pages are allocated. + indices = allocator.alloc_extend( + prefix_cpu.to(device), + prefix_cpu, + seq_cpu.to(device), + seq_cpu, + torch.tensor([-1], dtype=torch.int64, device=device), + 2 * ps, + ) + self.assertIsNotNone(indices) + decoded = allocator.alloc_decode( + (seq_cpu + 1).to(device), seq_cpu + 1, indices[-1:] + ) + self.assertIsNotNone(decoded) + indices = torch.cat((indices, decoded)) + self.assertTrue(torch.all(allocator.full_to_swa_index_mapping[indices] == 0)) + self.assertEqual(allocator.full_available_size(), ps) + return allocator, indices + + def test_swa_only_frees_leave_the_paged_pool_untouched(self): + for segment in (False, True): + for grouped in (False, True): + with self.subTest(segment=segment, grouped=grouped): + allocator, indices = self._allocated_ring() + swa_pages = ( + allocator.swa_attn_allocator.get_all_free_pages().clone() + ) + swa_available = allocator.swa_available_size() + if grouped: + allocator.free_group_begin() + if segment: + allocator.free_swa_segment(indices, start_pos=0) + else: + allocator.free_swa(indices) + self.assertEqual(allocator.swa_free_group, []) + self.assertEqual(allocator.swa_page_ids_group, []) + if grouped: + allocator.free_group_end() + self.assertTrue( + torch.equal( + allocator.swa_attn_allocator.get_all_free_pages(), swa_pages + ) + ) + self.assertEqual(allocator.swa_available_size(), swa_available) + self.assertEqual(allocator.full_available_size(), self.PS) + self.assertTrue( + torch.all(allocator.full_to_swa_index_mapping[indices] == 0) + ) + + def test_combined_frees_still_release_full_pages(self): + for segment in (False, True): + for grouped in (False, True): + with self.subTest(segment=segment, grouped=grouped): + allocator, indices = self._allocated_ring() + swa_pages = ( + allocator.swa_attn_allocator.get_all_free_pages().clone() + ) + if grouped: + allocator.free_group_begin() + if segment: + allocator.free_segment(indices, start_pos=0) + else: + allocator.free(indices) + if grouped: + self.assertEqual(allocator.full_available_size(), self.PS) + allocator.free_group_end() + self.assertEqual( + allocator.full_available_size(), allocator.size_full + ) + self.assertTrue( + torch.equal( + allocator.swa_attn_allocator.get_all_free_pages(), swa_pages + ) + ) + full_pages = allocator.full_attn_allocator.get_all_free_pages() + self.assertTrue(torch.all(full_pages > 0)) + self.assertEqual(torch.unique(full_pages).numel(), 4) + + def test_swa_only_frees_do_not_synchronize(self): + allocator, indices = self._allocated_ring() + peers = allocator.full_to_swa_index_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 envs.SGLANG_INVARIANT_CHECK.override(int(InvariantCheckLevel.STRICT)): + for grouped in (False, True): + with self.subTest(grouped=grouped): + if grouped: + allocator.free_group_begin() + self.assertIsNone(_sync_error(lambda: allocator.free_swa(indices))) + self.assertIsNone( + _sync_error( + lambda: allocator.free_swa_segment(indices, start_pos=0) + ) + ) + if grouped: + self.assertIsNone(_sync_error(allocator.free_group_end)) + + class TestSWAPageRepsFree(CustomTestCase): """page_size > 1: with a start position the SWA side frees one representative per page instead of expanding, filtering and dedup'ing through torch.unique.""" diff --git a/test/registered/unit/mem_cache/test_unified_cache_linker.py b/test/registered/unit/mem_cache/test_unified_cache_linker.py index fa48f6638..89c671a6c 100644 --- a/test/registered/unit/mem_cache/test_unified_cache_linker.py +++ b/test/registered/unit/mem_cache/test_unified_cache_linker.py @@ -4,6 +4,7 @@ from array import array from collections import defaultdict from dataclasses import replace from types import SimpleNamespace +from unittest.mock import MagicMock import pytest import test_unified_radix_cache_unittest as shared_cache_suite @@ -17,10 +18,13 @@ from test_unified_radix_cache_unittest import ( build_fixture, ) +from sglang.srt.managers.schedule_batch import ReqKvInfo +from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( InitLoadBackParams, InsertResult, MatchPrefixParams, + MatchResult, ) from sglang.srt.mem_cache.hicache_storage import ( PoolHitPolicy, @@ -36,8 +40,10 @@ from sglang.srt.mem_cache.unified_cache.components.full_component import FullCom from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent from sglang.srt.mem_cache.unified_cache.components.tree_component import ( ExternalLinkerLoadPhase, + LinkerTransferPhase, ) from sglang.srt.mem_cache.unified_cache.unified_cache_linker import ( + ExternalCacheHitMarker, UnifiedCacheLinker, UnifiedCacheLinkerWrapper, ) @@ -139,6 +145,8 @@ class _FakeExternalTreeCore: def _cache_for_wrapper(**kwargs): defaults = { + "_components_tuple": (), + "components": {}, "tree_core": SimpleNamespace(enable_external_cache_linker=False), "tree_components": (ComponentType.FULL,), "write_through_threshold": 256, @@ -149,6 +157,14 @@ def _cache_for_wrapper(**kwargs): return SimpleNamespace(**defaults) +def _swa_allocator(swa_req_ring): + if swa_req_ring is None: + return SimpleNamespace() + allocator = SWATokenToKVPoolAllocator.__new__(SWATokenToKVPoolAllocator) + allocator._swa_req_ring = swa_req_ring + return allocator + + def test_cache_linker_attachment_is_backend_independent(): cache = UnifiedRadixCache.__new__(UnifiedRadixCache) cache.tree_core = SimpleNamespace( @@ -157,6 +173,8 @@ def test_cache_linker_attachment_is_backend_independent(): ) cache.tree_components = (ComponentType.FULL,) cache.linker = None + cache._components_tuple = () + cache.components = {} linker = _FakeLinker() cache.init_cache_linker(linker) @@ -1063,5 +1081,257 @@ def test_component_commit_keeps_only_adopted_pages(): assert mapped_swa.tolist() == [202, 203, 206, 207] +@pytest.mark.parametrize( + "swa_req_ring", + [None, False, True], + ids=["other-allocator", "paged", "request-ring"], +) +@pytest.mark.parametrize("enable_hicache", [False, True]) +def test_swa_reuse_policy_tracks_layout_without_a_tier_condition( + monkeypatch, swa_req_ring, enable_hicache +): + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import env_gate + + unified_kv = swa_req_ring is True + monkeypatch.setattr(env_gate, "is_unified_kv_triton", lambda: unified_kv) + component = SWAComponent.__new__(SWAComponent) + component.sliding_window_size = 128 + cache = UnifiedRadixCache.__new__(UnifiedRadixCache) + cache.token_to_kv_pool_allocator = _swa_allocator(swa_req_ring) + cache.components = {ComponentType.SWA: component} + cache.cache_controller = object() if enable_hicache else None + cache.tree_core = SimpleNamespace( + enable_hicache=enable_hicache, + has_swa_host_pool=enable_hicache and not unified_kv, + ) + component.cache = cache + component.tree_core = cache.tree_core + # #32759: request-relative SWA needs tail re-prefill even without HiCache. + assert cache.swa_reprefill_tail_tokens() == (128 if unified_kv else 0) + node = SimpleNamespace( + component_data={ + ComponentType.SWA: SimpleNamespace(value=None, host_value=None) + }, + backuped=False, + evicted=False, + ) + assert component.create_match_validator(match_device_only=True)(node) is unified_kv + + +def test_cache_without_swa_needs_no_reprefill(): + cache = UnifiedRadixCache.__new__(UnifiedRadixCache) + cache.components = {} + assert cache.swa_reprefill_tail_tokens() == 0 + + +@pytest.fixture +def full_linker_component(): + def build_transfer(phase, node, keys): + keys = ["offload"] if phase == LinkerTransferPhase.OFFLOAD else list(keys) + return PoolTransfer( + name=PoolName.KV, + keys=keys, + device_indices=None + if phase == LinkerTransferPhase.LOOKUP + else torch.arange(len(keys) * 2), + ) + + return SimpleNamespace( + component_type=ComponentType.FULL, + build_external_linker_transfer=MagicMock(side_effect=build_transfer), + update_external_linker_load=lambda phase, req, full_transfer, transfer, prefix_len, **kwargs: ( + transfer + ), + ) + + +def test_linker_filters_request_relative_swa_from_lookup( + full_linker_component, +): + full = full_linker_component + swa = SWAComponent.__new__(SWAComponent) + swa.build_external_linker_transfer = MagicMock( + side_effect=AssertionError("excluded SWA reached linker") + ) + node = SimpleNamespace(id=1, external_cache_stored=False) + cache = _cache_for_wrapper( + _components_tuple=(full, swa), + components={ComponentType.FULL: full, ComponentType.SWA: swa}, + token_to_kv_pool_allocator=_swa_allocator(True), + tree_core=SimpleNamespace(enable_external_cache_linker=False, is_eagle=False), + page_size=2, + _all_reduce_attn_groups=lambda value, op: None, + get_last_hash_value=lambda node: None, + resolve_node_handle=lambda node_id: node, + inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=lambda: object()), + dec_lock_ref=MagicMock(), + ) + backend = _FakeLinker() + backend.restorable = [2] + wrapper = UnifiedCacheLinkerWrapper(cache, backend) + assert wrapper._components == (full,) + result = MatchResult( + device_indices=torch.empty(0, dtype=torch.int64), + last_device_node=0, + last_host_node=0, + best_match_node=0, + ) + matched = wrapper.match( + RadixKey(array("q", [1, 2, 3, 4])), SimpleNamespace(rid="match"), result + ) + assert matched.host_hit_length == 4 + assert [c.args[0] for c in full.build_external_linker_transfer.call_args_list] == [ + LinkerTransferPhase.LOOKUP, + ] + swa.build_external_linker_transfer.assert_not_called() + + +@pytest.mark.parametrize("swa_req_ring", [False, True], ids=["paged", "request-ring"]) +@pytest.mark.parametrize( + "completion", [True, False, None], ids=["success", "failure", "reset"] +) +def test_offload_filters_tree_core_swa_transfers_and_preserves_lifecycle( + swa_req_ring, completion +): + node = SimpleNamespace( + id=7, external_cache_stored=False, write_through_pending_id=None + ) + transfers = [ + PoolTransfer(name=PoolName.KV, keys=["page"]), + PoolTransfer(name=PoolName.SWA, keys=["page"]), + ] + core = _FakeExternalTreeCore({node.id: node}, transfers) + swa = SWAComponent.__new__(SWAComponent) + lock_params = object() + cache = _cache_for_wrapper( + components={ComponentType.SWA: swa}, + _components_tuple=(swa,), + token_to_kv_pool_allocator=_swa_allocator(swa_req_ring), + tree_core=core, + inc_lock_ref=MagicMock( + return_value=SimpleNamespace(to_dec_params=lambda: lock_params) + ), + dec_lock_ref=MagicMock(), + ) + backend = _FakeLinker() + wrapper = UnifiedCacheLinkerWrapper(cache, backend) + + wrapper.offload_nodes([node.id, node.id]) + expected = transfers[:1] if swa_req_ring else transfers + assert backend.queued_offloads == [expected] + assert core.offload_transfers == transfers + assert node.write_through_pending_id == node.id + assert not node.external_cache_stored + cache.inc_lock_ref.assert_called_once_with(node.id) + cache.dec_lock_ref.assert_not_called() + + if completion is None: + wrapper.reset() + assert backend.reset_count == 1 + else: + backend.completed_offloads.append(completion) + wrapper.commit_completed_offloads(wrapper.take_completed_offloads(1)) + assert not wrapper.pending_offloads + assert node.write_through_pending_id is None + assert node.external_cache_stored is (completion is True) + cache.dec_lock_ref.assert_called_once_with(node.id, lock_params) + if completion is True: + wrapper.offload_nodes([node.id]) + assert backend.queued_offloads == [expected] + else: + wrapper.offload_nodes([node.id]) + assert backend.queued_offloads == [expected, expected] + wrapper.reset() + + +@pytest.mark.parametrize( + "swa_req_ring,previous_boundary,expected_boundary", + [ + pytest.param(True, None, 4, id="unified-tombstones"), + pytest.param(True, 8, 8, id="preserve-existing-boundary"), + pytest.param(False, None, 2, id="paged-prepare-boundary"), + pytest.param(None, None, 2, id="other-allocator-prepare-boundary"), + ], +) +def test_linker_load_preserves_swa_boundaries( + full_linker_component, swa_req_ring, previous_boundary, expected_boundary +): + full = full_linker_component + swa = SWAComponent.__new__(SWAComponent) + participates = not swa_req_ring + + def prepare(phase, req, full_transfer, transfer, prefix_len, **kwargs): + if phase == ExternalLinkerLoadPhase.PREPARE: + req.kv = ReqKvInfo( + kv_allocated_len=prefix_len, swa_evicted_seqlen=prefix_len - 2 + ) + return transfer + + swa.build_external_linker_transfer = MagicMock( + return_value=PoolTransfer( + name=PoolName.SWA, keys=["a", "b"], device_indices=torch.arange(20, 24) + ) + ) + swa.update_external_linker_load = MagicMock(side_effect=prepare) + full_indices = torch.arange(4, dtype=torch.int64) + adopted = {ComponentType.FULL: [(0, 4)]} + if participates: + adopted[ComponentType.SWA] = [(0, 4)] + cache = _cache_for_wrapper( + _components_tuple=(full, swa), + page_size=2, + components={ComponentType.FULL: full, ComponentType.SWA: swa}, + token_to_kv_pool_allocator=_swa_allocator(swa_req_ring), + tree_core=SimpleNamespace( + empty_match_result=SimpleNamespace( + device_indices=torch.empty(0, dtype=torch.int64) + ), + collect_full_device_indices=lambda node, ancestor: full_indices, + mark_external_cache_stored_path=MagicMock(), + ), + insert=MagicMock( + return_value=InsertResult( + prefix_len=4, total_len=4, last_device_node=0, adopted_ranges=adopted + ) + ), + resolve_node_handle=lambda node_id: SimpleNamespace(id=0), + ) + wrapper = UnifiedCacheLinkerWrapper(cache, _FakeLinker()) + wrapper.hit_markers["rid"] = ExternalCacheHitMarker( + prefix_key=RadixKey(array("q", [1, 2, 3, 4])), + tail_hashes=["a", "b"], + device_hit_len=0, + ) + wrapper._queue_load = MagicMock() + kv = ( + None + if previous_boundary is None + else ReqKvInfo( + kv_allocated_len=previous_boundary, swa_evicted_seqlen=previous_boundary + ) + ) + req = SimpleNamespace( + rid="rid", + kv=kv, + prefix_indices=torch.empty(0, dtype=torch.int64), + last_node=0, + priority=0, + ) + restored, last_node = wrapper.load_back(req) + + assert restored.tolist() == full_indices.tolist() + assert last_node == 0 + assert req.kv.swa_evicted_seqlen == expected_boundary + assert req.kv.kv_allocated_len == (previous_boundary or 4) + assert cache.insert.call_args.args[0].swa_evicted_seqlen == expected_boundary + cache.tree_core.mark_external_cache_stored_path.assert_called_once_with(0, 0) + assert [c.args[0] for c in full.build_external_linker_transfer.call_args_list] == [ + LinkerTransferPhase.LOAD + ] + if not participates: + swa.build_external_linker_transfer.assert_not_called() + swa.update_external_linker_load.assert_not_called() + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__, "-v"]))