Files
sglang/test/registered/unit/mem_cache/test_linker_pool_assembler.py
T

292 lines
10 KiB
Python

"""Unit tests for external-linker device pool assembly."""
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.mem_cache.hicache_storage import (
PoolHitPolicy,
PoolName,
PoolTransfer,
)
from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
DevicePoolEntry,
DevicePoolGroup,
resolve_hybrid_device_pool_group,
)
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
class TestDevicePoolEntry(CustomTestCase):
def test_sparse_multi_component_layer_ranges(self):
k0 = torch.zeros((8, 3), dtype=torch.uint8)
k2 = torch.zeros((8, 5), dtype=torch.uint8)
v0 = torch.zeros((8, 7), dtype=torch.uint8)
v2 = torch.zeros((8, 11), dtype=torch.uint8)
pool = DevicePoolEntry(
name=PoolName.KV,
indices_from_pool=PoolName.KV,
device_pool=None,
components=[[k0, k2], [v0, v2]],
layer_mapping={0: 0, 2: 1},
page_size=2,
rows_are_pages=False,
packed=False,
)
indices = torch.tensor([0, 1, 4, 5])
locations = pool.prepare_locations(indices)
self.assertEqual(locations, [0, 4])
pointers, sizes = pool.get_page_buffer_meta(indices)
self.assertEqual(
pointers,
[
buffer[row].data_ptr()
for row in locations
for buffer in (k0, k2, v0, v2)
],
)
self.assertEqual(sizes, [6, 10, 14, 22] * 2)
self.assertIsNone(pool.get_prepared_layer_range_meta(locations, 1))
pointers, sizes, offsets = pool.get_prepared_layer_range_meta(locations, 2)
self.assertEqual(
pointers,
[
[k2[0].data_ptr()],
[v2[0].data_ptr()],
[k2[4].data_ptr()],
[v2[4].data_ptr()],
],
)
self.assertEqual(sizes, [[10], [22], [10], [22]])
self.assertEqual(offsets, [[6], [14], [6], [14]])
def test_rejects_invalid_pages_and_empty_buffers(self):
with self.assertRaisesRegex(ValueError, "has no storage buffers"):
DevicePoolEntry(
name=PoolName.KV,
indices_from_pool=PoolName.KV,
device_pool=None,
components=[],
layer_mapping={},
page_size=2,
rows_are_pages=False,
)
pool = DevicePoolEntry(
name=PoolName.KV,
indices_from_pool=PoolName.KV,
device_pool=None,
components=[[torch.zeros((8, 3), dtype=torch.uint8)]],
layer_mapping={0: 0},
page_size=2,
rows_are_pages=False,
)
for indices, error in (
(torch.tensor([0]), "multiple of page_size"),
(torch.tensor([1, 2]), "aligned contiguous pages"),
(torch.tensor([0, 2]), "aligned contiguous pages"),
(torch.tensor([8, 9]), "exceeds buffer shapes"),
):
with self.subTest(indices=indices.tolist()):
with self.assertRaisesRegex(ValueError, error):
pool.prepare_locations(indices)
class TestDevicePoolGroup(CustomTestCase):
def test_resolve_transfers_expands_physical_pools(self):
entries = [
SimpleNamespace(
name=PoolName.KV,
indices_from_pool=PoolName.KV,
translate_indices=lambda indices: indices,
),
SimpleNamespace(
name=PoolName.INDEXER,
indices_from_pool=PoolName.KV,
translate_indices=lambda indices: indices + 100,
),
]
group = DevicePoolGroup(entries, num_layers=2, page_size=2)
transfer = PoolTransfer(
name=PoolName.KV,
keys=["a", "b"],
device_indices=torch.tensor([0, 1, 4, 5]),
hit_policy=PoolHitPolicy.TRAILING_PAGES,
)
resolved = group.resolve_transfers([transfer])
self.assertEqual(
[item.name for item in resolved], [PoolName.KV, PoolName.INDEXER]
)
self.assertEqual(resolved[0].host_indices.tolist(), [0, 1, 4, 5])
self.assertEqual(resolved[1].host_indices.tolist(), [100, 101, 104, 105])
self.assertTrue(
all(item.hit_policy == PoolHitPolicy.ALL_PAGES for item in resolved)
)
def test_partial_side_pool_requires_explicit_opt_in(self):
entry = SimpleNamespace(
name=PoolName.SWA,
indices_from_pool=PoolName.SWA,
translate_indices=lambda indices: indices + 100,
)
group = DevicePoolGroup([entry], num_layers=1, page_size=2)
transfer = PoolTransfer(
name=PoolName.SWA,
keys=["b", "d"],
device_indices=torch.tensor([20, 21, 24, 25]),
hit_policy=PoolHitPolicy.TRAILING_PAGES,
)
self.assertEqual(group.resolve_transfers([transfer]), [])
resolved = group.resolve_transfers(
[transfer], allow_partial=True, allow_missing_kv=True
)
self.assertEqual(len(resolved), 1)
self.assertEqual(resolved[0].name, PoolName.SWA)
self.assertEqual(resolved[0].keys, ["b", "d"])
self.assertEqual(resolved[0].host_indices.tolist(), [120, 121, 124, 125])
self.assertEqual(resolved[0].hit_policy, PoolHitPolicy.TRAILING_PAGES)
class TestHybridDevicePoolAssembler(CustomTestCase):
def test_deepseek_v4_maps_sparse_sidecars(self):
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
DeepSeekV4LayerItem,
DeepSeekV4TokenToKVPool,
)
def state_pool():
return SimpleNamespace(
ring_size=2,
kv_score_buffer=SimpleNamespace(kv_score=torch.zeros((8, 3))),
)
kvcache = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
kvcache._unified_kv = False
kvcache.start_layer = 1
kvcache.end_layer = 4
kvcache.swa_page_size = 2
kvcache.swa_kv_pool = SimpleNamespace(
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)]
)
kvcache.c4_indexer_kv_pool = SimpleNamespace(
index_k_with_scale_buffer=[
torch.zeros((8, 7), dtype=torch.uint8) for _ in range(2)
]
)
kvcache.c128_kv_pool = SimpleNamespace(
kv_buffer=[torch.zeros((8, 11), dtype=torch.uint8)]
)
kvcache.layer_mapping = [
DeepSeekV4LayerItem(0, -1),
DeepSeekV4LayerItem(4, 0),
DeepSeekV4LayerItem(128, 0),
DeepSeekV4LayerItem(4, 1),
]
kvcache.compress_state_pools = [None, state_pool(), None, state_pool()]
kvcache.indexer_compress_state_pools = [
None,
state_pool(),
None,
state_pool(),
]
group = resolve_hybrid_device_pool_group(
kvcache=kvcache,
page_size=2,
params=SimpleNamespace(),
components={ComponentType.FULL, ComponentType.SWA},
)
self.assertEqual(group.num_layers, 3)
self.assertTrue(group.rank_replicated)
self.assertEqual(
set(group.entry_map),
{
PoolName.SWA,
PoolName.DEEPSEEK_V4_C4,
PoolName.DEEPSEEK_V4_C4_INDEXER,
PoolName.DEEPSEEK_V4_C128,
PoolName.DEEPSEEK_V4_C4_STATE,
PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
},
)
self.assertEqual(group.sources[PoolName.DEEPSEEK_V4_C4], PoolName.KV)
self.assertEqual(group.sources[PoolName.DEEPSEEK_V4_C4_STATE], PoolName.SWA)
c4_pool = group.entry_map[PoolName.DEEPSEEK_V4_C4]
pointers, sizes = c4_pool.get_page_buffer_meta(torch.tensor([0, 1]))
self.assertEqual(len(pointers), 2)
self.assertEqual(sizes, [5, 5])
_, sizes, offsets = c4_pool.get_prepared_layer_range_meta([0], 2)
self.assertEqual(sizes, [[5]])
self.assertEqual(offsets, [[5]])
self.assertIsNone(c4_pool.get_prepared_layer_range_meta([0], 1))
def test_dsa_uses_hybrid_assembler_strategy(self):
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
kvcache = DSATokenToKVPool.__new__(DSATokenToKVPool)
kvcache.page_size = 2
kvcache.layer_num = 2
kvcache.kv_buffer = [
torch.zeros((8, 3), dtype=torch.uint8),
torch.zeros((8, 5), dtype=torch.uint8),
]
kvcache.index_key_cache = SimpleNamespace(
buffer=[
torch.zeros((4, 7), dtype=torch.uint8),
torch.zeros((4, 11), dtype=torch.uint8),
]
)
group = resolve_hybrid_device_pool_group(
kvcache=kvcache,
page_size=2,
params=SimpleNamespace(),
components={ComponentType.FULL},
)
self.assertEqual(group.num_layers, 2)
self.assertTrue(group.rank_replicated)
self.assertEqual(set(group.entry_map), {PoolName.KV, PoolName.INDEXER})
self.assertEqual(
group.sources,
{
PoolName.KV: PoolName.KV,
PoolName.INDEXER: PoolName.KV,
},
)
def test_unsupported_strategy_fails_with_context(self):
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
kvcache = HybridLinearKVPool.__new__(HybridLinearKVPool)
with self.assertRaisesRegex(
ValueError,
"does not support the direct external linker: _MambaStrategy",
):
resolve_hybrid_device_pool_group(
kvcache=kvcache,
page_size=2,
params=SimpleNamespace(),
components={ComponentType.FULL, ComponentType.MAMBA},
)
if __name__ == "__main__":
unittest.main()