[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):
"""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"]))