[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:
co-authored by
amd-danli103
Duyi-Wang
TianDi101
parent
0bae67648a
commit
822e73ccdd
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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"]))
|
||||
|
||||
Reference in New Issue
Block a user