[Unified Cache][AMD] Support DeepSeek-V4 unified KV in direct external linkers (#38269)

Co-authored-by: amd-danli103 <danli103@amd.com>
Co-authored-by: Duyi-Wang <duyi.wang@amd.com>
Co-authored-by: TianDi101 <tiandi920722@gmail.com>
This commit is contained in:
Niko Ma
2026-09-11 01:49:29 -07:00
committed by GitHub
co-authored by amd-danli103 Duyi-Wang TianDi101
parent 0bae67648a
commit 822e73ccdd
7 changed files with 572 additions and 56 deletions
+6 -4
View File
@@ -431,8 +431,9 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def free_swa(self, free_index: torch.Tensor): def free_swa(self, free_index: torch.Tensor):
"""Release the SWA peers of an arbitrary slot set and clear their mapping. """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().""" No-op for a per-request ring, which owns no paged SWA peers. Otherwise
if free_index.numel() == 0: synchronizes at page_size > 1; kv-row segments use free_swa_segment()."""
if self._swa_req_ring or free_index.numel() == 0:
return return
if self.page_size == 1: if self.page_size == 1:
@@ -455,8 +456,9 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def free_swa_segment(self, free_index: torch.Tensor, *, start_pos: int): 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_swa() for a kv-row segment; same start-alignment contract as
free_segment(), and fixed-shape at every page size.""" free_segment(), and fixed-shape at every page size. No-op for a
if free_index.numel() == 0: per-request ring, as in free_swa()."""
if self._swa_req_ring or free_index.numel() == 0:
return return
self._free_swa_pages(free_index, start_pos=start_pos) self._free_swa_pages(free_index, start_pos=start_pos)
@@ -239,24 +239,24 @@ def _build_deepseek_v4_device_pool_group(
) -> DevicePoolGroup: ) -> DevicePoolGroup:
from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
_dsv4_compressed_region_buffers,
_dsv4_indexer_regions, _dsv4_indexer_regions,
_resolve_deepseek_v4_layer_mappings, _resolve_deepseek_v4_layer_mappings,
) )
if isinstance(kvcache.c4_kv_pool, HiSparseC4DevicePool):
raise ValueError("The direct external linker does not support HiSparse.")
mappings = _resolve_deepseek_v4_layer_mappings(kvcache) mappings = _resolve_deepseek_v4_layer_mappings(kvcache)
if getattr(kvcache, "_unified_kv", False) or isinstance( is_unified_kv = getattr(kvcache, "_unified_kv", False)
kvcache.c4_kv_pool, HiSparseC4DevicePool entries = []
): if not is_unified_kv:
raise ValueError(
"The direct external linker does not support unified-KV or HiSparse."
)
if kvcache.swa_page_size != page_size: if kvcache.swa_page_size != page_size:
raise ValueError( raise ValueError(
"DeepSeek V4 SWA page size must match the tree page size: " "DeepSeek V4 SWA page size must match the tree page size: "
f"{kvcache.swa_page_size} != {page_size}." f"{kvcache.swa_page_size} != {page_size}."
) )
entries.append(
entries = [
DevicePoolEntry( DevicePoolEntry(
name=PoolName.SWA, name=PoolName.SWA,
indices_from_pool=PoolName.SWA, indices_from_pool=PoolName.SWA,
@@ -266,7 +266,7 @@ def _build_deepseek_v4_device_pool_group(
page_size=page_size, page_size=page_size,
rows_are_pages=True, rows_are_pages=True,
) )
] )
def add(name, source, pool, buffers, layer_mapping): def add(name, source, pool, buffers, layer_mapping):
if 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( add(
PoolName.DEEPSEEK_V4_C4, PoolName.DEEPSEEK_V4_C4,
PoolName.KV, PoolName.KV,
kvcache.c4_kv_pool, kvcache.c4_kv_pool,
kvcache.c4_kv_pool.kv_buffer, c4_buffers,
mappings.c4, mappings.c4,
) )
for region in _dsv4_indexer_regions(kvcache, page_size): for region in _dsv4_indexer_regions(kvcache, page_size):
@@ -301,9 +304,10 @@ def _build_deepseek_v4_device_pool_group(
PoolName.DEEPSEEK_V4_C128, PoolName.DEEPSEEK_V4_C128,
PoolName.KV, PoolName.KV,
kvcache.c128_kv_pool, kvcache.c128_kv_pool,
kvcache.c128_kv_pool.kv_buffer, c128_buffers,
mappings.c128, mappings.c128,
) )
if not is_unified_kv:
add( add(
PoolName.DEEPSEEK_V4_C4_STATE, PoolName.DEEPSEEK_V4_C4_STATE,
PoolName.SWA, PoolName.SWA,
@@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Callable, Optional, Sequence
import torch import torch
from sglang.srt.environ import envs 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 ( from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams, DecLockRefParams,
EvictParams, EvictParams,
@@ -308,11 +309,9 @@ class SWAComponent(TreeComponent):
ct = self.component_type ct = self.component_type
state = {"len": float("inf")} state = {"len": float("inf")}
# unified_kv never caches the SWA ring (per-request, not content-stable), # A per-request SWA ring is not stored in tree nodes, so its bookkeeping
# so SWA bookkeeping must not gate the match here. # must not gate prefix matching.
swa_device_only_hicache = ( swa_req_ring = is_swa_req_ring(self.cache.token_to_kv_pool_allocator)
not self.tree_core.has_swa_host_pool and self.tree_core.enable_hicache
)
def validator(node: UnifiedTreeNode) -> bool: def validator(node: UnifiedTreeNode) -> bool:
cd = node.component_data[ct] cd = node.component_data[ct]
@@ -320,7 +319,7 @@ class SWAComponent(TreeComponent):
# — load_back will restore SWA from host before use. # — load_back will restore SWA from host before use.
if cd.value is None and (match_device_only or cd.host_value is None): if cd.value is None and (match_device_only or cd.host_value is None):
state["len"] = 0 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 True
return False return False
state["len"] += len(node.key) state["len"] += len(node.key)
@@ -25,6 +25,7 @@ from typing import TYPE_CHECKING, NamedTuple
import torch import torch
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams, DecLockRefParams,
InsertParams, InsertParams,
@@ -160,6 +161,15 @@ class UnifiedCacheLinkerWrapper:
self.cache = cache self.cache = cache
self.cache_linker = cache_linker 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. # rid -> what match found, consumed by the next init_load_back.
self.hit_markers: dict[str, ExternalCacheHitMarker] = {} self.hit_markers: dict[str, ExternalCacheHitMarker] = {}
# Loads in flight, each pinning its inserted endpoint until DMA completes. # Loads in flight, each pinning its inserted endpoint until DMA completes.
@@ -192,7 +202,7 @@ class UnifiedCacheLinkerWrapper:
return result return result
lookup_transfers = [] lookup_transfers = []
for component in cache._components_tuple: for component in self._components:
transfer = component.build_external_linker_transfer( transfer = component.build_external_linker_transfer(
LinkerTransferPhase.LOOKUP, None, tail_hashes LinkerTransferPhase.LOOKUP, None, tail_hashes
) )
@@ -290,7 +300,7 @@ class UnifiedCacheLinkerWrapper:
# Build per-component linker transfers. # Build per-component linker transfers.
component_transfers: list[tuple[TreeComponent, PoolTransfer]] = [] component_transfers: list[tuple[TreeComponent, PoolTransfer]] = []
for component in cache._components_tuple: for component in self._components:
transfer = component.build_external_linker_transfer( transfer = component.build_external_linker_transfer(
LinkerTransferPhase.LOAD, None, tail_hashes LinkerTransferPhase.LOAD, None, tail_hashes
) )
@@ -313,6 +323,20 @@ class UnifiedCacheLinkerWrapper:
prefix_len, 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. # Insert the newly loaded tail into the tree.
prefix_indices = torch.cat( prefix_indices = torch.cat(
[req.prefix_indices.to(torch.int64), full_transfer.device_indices] [req.prefix_indices.to(torch.int64), full_transfer.device_indices]
@@ -484,6 +508,8 @@ class UnifiedCacheLinkerWrapper:
node_id node_id
) )
if transfers is not None: 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) self._offload_node(node_id, transfers)
def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None: def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None:
@@ -2,6 +2,7 @@
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import Mock, call
import torch 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 ( from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
DevicePoolEntry, DevicePoolEntry,
DevicePoolGroup, DevicePoolGroup,
_build_deepseek_v4_device_pool_group,
resolve_hybrid_device_pool_group, resolve_hybrid_device_pool_group,
) )
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType 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)] kv_buffer=[torch.zeros((8, 3), dtype=torch.uint8) for _ in range(3)]
) )
kvcache.c4_kv_pool = SimpleNamespace( 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( kvcache.c4_indexer_kv_pool = SimpleNamespace(
index_k_with_scale_buffer=[ index_k_with_scale_buffer=[
@@ -188,7 +191,8 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
] ]
) )
kvcache.c128_kv_pool = SimpleNamespace( 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 = [ kvcache.layer_mapping = [
DeepSeekV4LayerItem(0, -1), DeepSeekV4LayerItem(0, -1),
@@ -236,6 +240,95 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
self.assertEqual(offsets, [[5]]) self.assertEqual(offsets, [[5]])
self.assertIsNone(c4_pool.get_prepared_layer_range_meta([0], 1)) 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): def test_dsa_uses_hybrid_assembler_strategy(self):
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
@@ -58,6 +58,7 @@ def _build_swa_tree(
kv_size_swa: int = 32, kv_size_swa: int = 32,
sliding_window_size: int = 4, sliding_window_size: int = 4,
enable_kv_cache_events: bool = False, enable_kv_cache_events: bool = False,
swa_req_ring_size: int | None = None,
): ):
head_num = 8 head_num = 8
head_dim = 128 head_dim = 128
@@ -88,6 +89,7 @@ def _build_swa_tree(
full_attention_layer_ids=full_attention_layer_ids, full_attention_layer_ids=full_attention_layer_ids,
device=device, device=device,
) )
kv_pool.swa_req_ring_size = swa_req_ring_size
allocator = SWATokenToKVPoolAllocator( allocator = SWATokenToKVPoolAllocator(
size=kv_size, size=kv_size,
size_swa=kv_size_swa, size_swa=kv_size_swa,
@@ -96,6 +98,7 @@ def _build_swa_tree(
device=device, device=device,
kvcache=kv_pool, kvcache=kv_pool,
need_sort=False, need_sort=False,
req_to_token_pool=req_to_token_pool,
) )
tree = SWARadixCache( tree = SWARadixCache(
params=CacheInitParams( params=CacheInitParams(
@@ -1234,6 +1237,125 @@ class TestSWAPeerMappedContract(CustomTestCase):
self.assertIsNone(_sync_error(lambda: allocator.free_swa(indices))) 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): class TestSWAPageRepsFree(CustomTestCase):
"""page_size > 1: with a start position the SWA side frees one representative """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.""" per page instead of expanding, filtering and dedup'ing through torch.unique."""
@@ -4,6 +4,7 @@ from array import array
from collections import defaultdict from collections import defaultdict
from dataclasses import replace from dataclasses import replace
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest import pytest
import test_unified_radix_cache_unittest as shared_cache_suite import test_unified_radix_cache_unittest as shared_cache_suite
@@ -17,10 +18,13 @@ from test_unified_radix_cache_unittest import (
build_fixture, 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 ( from sglang.srt.mem_cache.base_prefix_cache import (
InitLoadBackParams, InitLoadBackParams,
InsertResult, InsertResult,
MatchPrefixParams, MatchPrefixParams,
MatchResult,
) )
from sglang.srt.mem_cache.hicache_storage import ( from sglang.srt.mem_cache.hicache_storage import (
PoolHitPolicy, 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.swa_component import SWAComponent
from sglang.srt.mem_cache.unified_cache.components.tree_component import ( from sglang.srt.mem_cache.unified_cache.components.tree_component import (
ExternalLinkerLoadPhase, ExternalLinkerLoadPhase,
LinkerTransferPhase,
) )
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import ( from sglang.srt.mem_cache.unified_cache.unified_cache_linker import (
ExternalCacheHitMarker,
UnifiedCacheLinker, UnifiedCacheLinker,
UnifiedCacheLinkerWrapper, UnifiedCacheLinkerWrapper,
) )
@@ -139,6 +145,8 @@ class _FakeExternalTreeCore:
def _cache_for_wrapper(**kwargs): def _cache_for_wrapper(**kwargs):
defaults = { defaults = {
"_components_tuple": (),
"components": {},
"tree_core": SimpleNamespace(enable_external_cache_linker=False), "tree_core": SimpleNamespace(enable_external_cache_linker=False),
"tree_components": (ComponentType.FULL,), "tree_components": (ComponentType.FULL,),
"write_through_threshold": 256, "write_through_threshold": 256,
@@ -149,6 +157,14 @@ def _cache_for_wrapper(**kwargs):
return SimpleNamespace(**defaults) 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(): def test_cache_linker_attachment_is_backend_independent():
cache = UnifiedRadixCache.__new__(UnifiedRadixCache) cache = UnifiedRadixCache.__new__(UnifiedRadixCache)
cache.tree_core = SimpleNamespace( cache.tree_core = SimpleNamespace(
@@ -157,6 +173,8 @@ def test_cache_linker_attachment_is_backend_independent():
) )
cache.tree_components = (ComponentType.FULL,) cache.tree_components = (ComponentType.FULL,)
cache.linker = None cache.linker = None
cache._components_tuple = ()
cache.components = {}
linker = _FakeLinker() linker = _FakeLinker()
cache.init_cache_linker(linker) 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] 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__": if __name__ == "__main__":
raise SystemExit(pytest.main([__file__, "-v"])) raise SystemExit(pytest.main([__file__, "-v"]))