sgl-project--sglang
94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
745 行
29 KiB
Python
745 行
29 KiB
Python
import unittest
|
|
from array import array
|
|
|
|
import torch
|
|
|
|
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
|
|
from sglang.srt.environ import envs
|
|
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
|
DecLockRefParams,
|
|
EvictParams,
|
|
EvictResult,
|
|
InsertParams,
|
|
MatchPrefixParams,
|
|
)
|
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
|
from sglang.srt.mem_cache.common import available_and_evictable_str
|
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
|
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
|
from sglang.srt.utils import get_device
|
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
|
from sglang.test.test_utils import CustomTestCase
|
|
|
|
register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-large")
|
|
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd")
|
|
|
|
|
|
def _event_hashes(events):
|
|
return [block_hash for event in events for block_hash in event.block_hashes]
|
|
|
|
|
|
class _DummyReq:
|
|
def __init__(self):
|
|
self._kv_committed_len = 0
|
|
self.swa_prefix_lock_released = False
|
|
|
|
def pop_committed_kv_cache(self):
|
|
return self._kv_committed_len
|
|
|
|
|
|
def _build_swa_tree(
|
|
is_eagle: bool,
|
|
page_size: int = 1,
|
|
req_size: int = 8,
|
|
max_context_len: int = 64,
|
|
kv_size: int = 64,
|
|
kv_size_swa: int = 32,
|
|
sliding_window_size: int = 4,
|
|
enable_kv_cache_events: bool = False,
|
|
):
|
|
head_num = 8
|
|
head_dim = 128
|
|
num_layers = 24
|
|
global_interval = 4
|
|
dtype = torch.bfloat16
|
|
device = get_device()
|
|
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
|
|
full_attention_layer_ids_set = set(full_attention_layer_ids)
|
|
swa_attention_layer_ids = [
|
|
i for i in range(num_layers) if i not in full_attention_layer_ids_set
|
|
]
|
|
|
|
req_to_token_pool = ReqToTokenPool(
|
|
size=req_size,
|
|
max_context_len=max_context_len,
|
|
device=device,
|
|
enable_memory_saver=False,
|
|
)
|
|
kv_pool = SWAKVPool(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
head_num=head_num,
|
|
head_dim=head_dim,
|
|
swa_attention_layer_ids=swa_attention_layer_ids,
|
|
full_attention_layer_ids=full_attention_layer_ids,
|
|
device=device,
|
|
)
|
|
allocator = SWATokenToKVPoolAllocator(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
device=device,
|
|
kvcache=kv_pool,
|
|
need_sort=False,
|
|
)
|
|
tree = SWARadixCache(
|
|
params=CacheInitParams(
|
|
req_to_token_pool=req_to_token_pool,
|
|
token_to_kv_pool_allocator=allocator,
|
|
page_size=page_size,
|
|
disable=False,
|
|
is_eagle=is_eagle,
|
|
sliding_window_size=sliding_window_size,
|
|
enable_kv_cache_events=enable_kv_cache_events,
|
|
),
|
|
)
|
|
return tree, allocator, req_to_token_pool
|
|
|
|
|
|
def _swa_alloc(allocator, need_size):
|
|
"""SWA-pool alloc that also works for page_size > 1 (built-in alloc asserts page_size == 1)."""
|
|
if allocator.page_size == 1:
|
|
return allocator.alloc(need_size)
|
|
|
|
assert need_size % allocator.page_size == 0
|
|
full_indices = allocator.full_attn_allocator.alloc(need_size)
|
|
swa_indices = allocator.swa_attn_allocator.alloc(need_size)
|
|
assert full_indices is not None and swa_indices is not None
|
|
allocator.full_to_swa_index_mapping[full_indices] = swa_indices
|
|
return full_indices
|
|
|
|
|
|
def _insert(tree, allocator, token_ids):
|
|
indices = _swa_alloc(allocator, len(token_ids))
|
|
assert indices is not None
|
|
tree.insert(InsertParams(key=RadixKey(array("q", token_ids)), value=indices))
|
|
|
|
|
|
def _insert_chain(tree, allocator, token_ids):
|
|
_insert(tree, allocator, token_ids)
|
|
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", token_ids))))
|
|
return match.last_device_node
|
|
|
|
|
|
def _expected_tail_size(window: int, page_size: int) -> int:
|
|
"""Mirror of _maybe_split_leaf_for_swa_lock's tail_size formula."""
|
|
return (window + page_size - 1) // page_size * page_size
|
|
|
|
|
|
class TestSWA(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
pass
|
|
|
|
@classmethod
|
|
def tearDownClass(cls):
|
|
pass
|
|
|
|
def test_swa_radix_cache_kv_events(self):
|
|
tree, allocator, _ = _build_swa_tree(
|
|
is_eagle=False, enable_kv_cache_events=True
|
|
)
|
|
tree.take_events() # Clear the reset event.
|
|
|
|
_insert(tree, allocator, [1, 2, 3, 4])
|
|
first_insert_events = [
|
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
|
]
|
|
self.assertEqual(len(first_insert_events), 4)
|
|
self.assertEqual([e.token_ids[0] for e in first_insert_events], [1, 2, 3, 4])
|
|
|
|
_insert(tree, allocator, [1, 2, 3, 4, 5, 6])
|
|
second_insert_events = [
|
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
|
]
|
|
self.assertEqual(len(second_insert_events), 2)
|
|
self.assertEqual([e.token_ids[0] for e in second_insert_events], [5, 6])
|
|
|
|
stored_hashes = [
|
|
e.block_hashes[0] for e in first_insert_events + second_insert_events
|
|
]
|
|
|
|
# Evicting only SWA tokens tombstones nodes but keeps full KV blocks.
|
|
result = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=1))
|
|
self.assertEqual(result.num_tokens_evicted, 0)
|
|
self.assertGreaterEqual(result.swa_num_tokens_evicted, 1)
|
|
self.assertEqual(
|
|
[e for e in tree.take_events() if isinstance(e, BlockRemoved)], []
|
|
)
|
|
|
|
result = tree.evict(EvictParams(num_tokens=1, swa_num_tokens=0))
|
|
self.assertGreaterEqual(result.num_tokens_evicted, 1)
|
|
removed_hashes = _event_hashes(
|
|
[e for e in tree.take_events() if isinstance(e, BlockRemoved)]
|
|
)
|
|
self.assertCountEqual(removed_hashes, stored_hashes)
|
|
|
|
def test_swa_radix_cache_kv_events_split_hash(self):
|
|
tree, allocator, _ = _build_swa_tree(
|
|
is_eagle=False, enable_kv_cache_events=True
|
|
)
|
|
tree.take_events() # Clear the reset event.
|
|
|
|
_insert(tree, allocator, [1, 2, 3, 4])
|
|
first_insert_events = [
|
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
|
]
|
|
self.assertEqual(len(first_insert_events), 4)
|
|
split_parent_hash = first_insert_events[1].block_hashes[0]
|
|
|
|
_insert(tree, allocator, [1, 2, 5, 6])
|
|
second_insert_events = [
|
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
|
]
|
|
self.assertEqual(len(second_insert_events), 2)
|
|
self.assertEqual(list(second_insert_events[0].token_ids), [5])
|
|
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
|
|
|
|
def test_swa_memory_pool_paged_free_clears_full_page_mapping(self):
|
|
page_size = 4
|
|
_, allocator, _ = _build_swa_tree(
|
|
is_eagle=False,
|
|
page_size=page_size,
|
|
kv_size=16,
|
|
kv_size_swa=16,
|
|
sliding_window_size=page_size,
|
|
)
|
|
|
|
full_indices = _swa_alloc(allocator, page_size)
|
|
self.assertEqual(allocator.swa_available_size(), 16 - page_size)
|
|
|
|
allocator.free_swa(full_indices[:1])
|
|
self.assertEqual(allocator.swa_available_size(), 16)
|
|
self.assertTrue(
|
|
torch.all(
|
|
allocator.full_to_swa_index_mapping[full_indices.to(torch.int64)] == 0
|
|
)
|
|
)
|
|
|
|
allocator.free_swa(full_indices[1:2])
|
|
self.assertEqual(allocator.swa_available_size(), 16)
|
|
|
|
def test_swa_radix_cache_1(self):
|
|
# args
|
|
req_size = 10
|
|
max_context_len = 128
|
|
kv_size = 128
|
|
kv_size_swa = 64
|
|
page_size = 1
|
|
sliding_window_size = 4
|
|
head_num = 8
|
|
head_dim = 128
|
|
num_layers = 48
|
|
global_interval = 4
|
|
dtype = torch.bfloat16
|
|
device = get_device()
|
|
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
|
|
full_attention_layer_ids_set = set(full_attention_layer_ids)
|
|
swa_attention_layer_ids = [
|
|
i for i in range(num_layers) if i not in full_attention_layer_ids_set
|
|
]
|
|
# setup req to token pool
|
|
req_to_token_pool = ReqToTokenPool(
|
|
size=req_size,
|
|
max_context_len=max_context_len,
|
|
device=device,
|
|
enable_memory_saver=False,
|
|
)
|
|
# setup kv pool
|
|
kv_pool = SWAKVPool(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
head_num=head_num,
|
|
head_dim=head_dim,
|
|
swa_attention_layer_ids=swa_attention_layer_ids,
|
|
full_attention_layer_ids=full_attention_layer_ids,
|
|
device=device,
|
|
)
|
|
# setup token to kv pool allocator
|
|
allocator = SWATokenToKVPoolAllocator(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
device=device,
|
|
kvcache=kv_pool,
|
|
need_sort=False,
|
|
)
|
|
# setup radix cache
|
|
tree = SWARadixCache(
|
|
params=CacheInitParams(
|
|
req_to_token_pool=req_to_token_pool,
|
|
token_to_kv_pool_allocator=allocator,
|
|
disable=False,
|
|
page_size=page_size,
|
|
sliding_window_size=sliding_window_size,
|
|
),
|
|
)
|
|
|
|
# test
|
|
print(
|
|
f"[Start] allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
|
|
self.assertEqual(len(req1_token_ids), len(req1_kv_indices))
|
|
print(
|
|
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req1_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
print(
|
|
f"req1: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
|
|
self.assertEqual(len(req2_token_ids), len(req2_kv_indices))
|
|
print(
|
|
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req2_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
print(
|
|
f"req2: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req3_token_ids, req3_kv_indices = [10, 11, 12], allocator.alloc(3)
|
|
self.assertEqual(len(req3_token_ids), len(req3_kv_indices))
|
|
print(
|
|
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req3_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
print(
|
|
f"req3: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
|
|
self.assertEqual(len(req4_token_ids), len(req4_kv_indices))
|
|
print(
|
|
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req4_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
print(
|
|
f"req4: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
|
|
tree.pretty_print()
|
|
full_num_tokens, swa_num_tokens = 1, 0
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
tree.pretty_print()
|
|
|
|
full_num_tokens, swa_num_tokens = 0, 1
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
tree.pretty_print()
|
|
|
|
full_num_tokens, swa_num_tokens = 1, 2
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
tree.pretty_print()
|
|
|
|
req5_token_ids = [1, 2, 3, 4, 5]
|
|
result = tree.match_prefix(
|
|
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
|
|
)
|
|
kv_indices, last_node = result.device_indices, result.last_device_node
|
|
print(
|
|
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
|
|
)
|
|
self.assertEqual(len(kv_indices), 0)
|
|
|
|
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
|
|
result = tree.match_prefix(
|
|
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
|
|
)
|
|
kv_indices, last_node = result.device_indices, result.last_device_node
|
|
print(
|
|
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
|
|
)
|
|
self.assertEqual(len(kv_indices), 7)
|
|
self.assertEqual(len(last_node.key), 2)
|
|
self.assertEqual(last_node.key.token_ids[0], 60)
|
|
self.assertEqual(last_node.key.token_ids[1], 70)
|
|
|
|
print(tree.available_and_evictable_str())
|
|
print(available_and_evictable_str(tree))
|
|
tree.sanity_check()
|
|
|
|
def test_swa_radix_cache_eagle(self):
|
|
# args
|
|
req_size = 10
|
|
max_context_len = 128
|
|
kv_size = 128
|
|
kv_size_swa = 64
|
|
page_size = 1
|
|
sliding_window_size = 4
|
|
head_num = 8
|
|
head_dim = 128
|
|
num_layers = 48
|
|
global_interval = 4
|
|
dtype = torch.bfloat16
|
|
device = get_device()
|
|
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
|
|
full_attention_layer_ids_set = set(full_attention_layer_ids)
|
|
swa_attention_layer_ids = [
|
|
i for i in range(num_layers) if i not in full_attention_layer_ids_set
|
|
]
|
|
# setup req to token pool
|
|
req_to_token_pool = ReqToTokenPool(
|
|
size=req_size,
|
|
max_context_len=max_context_len,
|
|
device=device,
|
|
enable_memory_saver=False,
|
|
)
|
|
# setup kv pool
|
|
kv_pool = SWAKVPool(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
head_num=head_num,
|
|
head_dim=head_dim,
|
|
swa_attention_layer_ids=swa_attention_layer_ids,
|
|
full_attention_layer_ids=full_attention_layer_ids,
|
|
device=device,
|
|
)
|
|
# setup token to kv pool allocator
|
|
allocator = SWATokenToKVPoolAllocator(
|
|
size=kv_size,
|
|
size_swa=kv_size_swa,
|
|
page_size=page_size,
|
|
dtype=dtype,
|
|
device=device,
|
|
kvcache=kv_pool,
|
|
need_sort=False,
|
|
)
|
|
# setup radix cache
|
|
tree = SWARadixCache(
|
|
params=CacheInitParams(
|
|
req_to_token_pool=req_to_token_pool,
|
|
token_to_kv_pool_allocator=allocator,
|
|
page_size=page_size,
|
|
disable=False,
|
|
is_eagle=True,
|
|
sliding_window_size=sliding_window_size,
|
|
),
|
|
)
|
|
|
|
# test
|
|
print(
|
|
f"[Start] allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
|
|
self.assertEqual(len(req1_token_ids), len(req1_kv_indices))
|
|
print(
|
|
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req1_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
self.assertEqual(prefix_len, 0)
|
|
print(
|
|
f"req1: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
|
|
self.assertEqual(len(req2_token_ids), len(req2_kv_indices))
|
|
print(
|
|
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req2_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
self.assertEqual(prefix_len, 2)
|
|
print(
|
|
f"req2: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req3_token_ids, req3_kv_indices = [10, 11, 12], allocator.alloc(3)
|
|
self.assertEqual(len(req3_token_ids), len(req3_kv_indices))
|
|
print(
|
|
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req3_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
self.assertEqual(prefix_len, 0)
|
|
print(
|
|
f"req3: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
|
|
self.assertEqual(len(req4_token_ids), len(req4_kv_indices))
|
|
print(
|
|
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
|
|
)
|
|
key = RadixKey(array("q", req4_token_ids))
|
|
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
|
|
prefix_len = result.prefix_len
|
|
self.assertEqual(prefix_len, 4)
|
|
print(
|
|
f"req4: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
|
|
)
|
|
|
|
tree.pretty_print()
|
|
full_num_tokens, swa_num_tokens = 1, 0
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
evict_result = tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
assert isinstance(evict_result, EvictResult)
|
|
assert (
|
|
evict_result.num_tokens_evicted >= full_num_tokens
|
|
) # May evict more due to node granularity
|
|
print(
|
|
f"evicted {evict_result.num_tokens_evicted} full tokens, {evict_result.swa_num_tokens_evicted} swa tokens"
|
|
)
|
|
tree.pretty_print()
|
|
|
|
full_num_tokens, swa_num_tokens = 0, 1
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
evict_result = tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
assert isinstance(evict_result, EvictResult)
|
|
assert (
|
|
evict_result.swa_num_tokens_evicted >= swa_num_tokens
|
|
), f"evicted {evict_result.swa_num_tokens_evicted} swa tokens, expected {swa_num_tokens}"
|
|
tree.pretty_print()
|
|
|
|
full_num_tokens, swa_num_tokens = 1, 2
|
|
print(f"evicting {full_num_tokens} full token and {swa_num_tokens} swa token")
|
|
evict_result = tree.evict(
|
|
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
)
|
|
assert isinstance(evict_result, EvictResult)
|
|
assert (
|
|
evict_result.num_tokens_evicted >= full_num_tokens
|
|
), f"evicted {evict_result.num_tokens_evicted} full tokens, expected {full_num_tokens}"
|
|
assert (
|
|
evict_result.swa_num_tokens_evicted >= swa_num_tokens
|
|
), f"evicted {evict_result.swa_num_tokens_evicted} swa tokens, expected {swa_num_tokens}"
|
|
tree.pretty_print()
|
|
|
|
req5_token_ids = [1, 2, 3, 4, 5]
|
|
result = tree.match_prefix(
|
|
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
|
|
)
|
|
kv_indices, last_node = result.device_indices, result.last_device_node
|
|
print(
|
|
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
|
|
)
|
|
self.assertEqual(len(kv_indices), 0) # no swa prefix matched
|
|
|
|
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
|
|
result = tree.match_prefix(
|
|
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
|
|
)
|
|
kv_indices, last_node = result.device_indices, result.last_device_node
|
|
print(
|
|
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
|
|
)
|
|
self.assertEqual(len(kv_indices), 6)
|
|
self.assertEqual(len(last_node.key), 2)
|
|
# Bigram view: token_ids holds raw tokens; iteration yields bigram tuples.
|
|
self.assertTrue(last_node.key.is_bigram)
|
|
self.assertEqual(list(last_node.key), [(5, 60), (60, 70)])
|
|
|
|
def test_swa_cache_finished_req_eagle_uses_cache_protected_len_and_bigram_key(self):
|
|
tree, allocator, req_to_token_pool = _build_swa_tree(is_eagle=True)
|
|
|
|
# Case 1: is_insert=True should pass bigram key and use cache_protected_len.
|
|
req = _DummyReq()
|
|
req.req_pool_idx = 0
|
|
req.origin_input_ids = array("q", [1, 2, 3, 4, 5, 6])
|
|
req.output_ids = array("q")
|
|
req._kv_committed_len = len(req.origin_input_ids)
|
|
kv_indices = allocator.alloc(req._kv_committed_len)
|
|
req_to_token_pool.write(
|
|
(req.req_pool_idx, slice(0, req._kv_committed_len)), kv_indices
|
|
)
|
|
req.extra_key = None
|
|
req.last_node = tree.root_node
|
|
req.swa_uuid_for_lock = None
|
|
req.swa_evicted_seqlen = 0
|
|
req.cache_protected_len = 1
|
|
# Intentionally mismatch to ensure code does not use len(prefix_indices).
|
|
req.prefix_indices = torch.tensor([7, 8, 9, 10, 11], device=tree.device)
|
|
|
|
captured = {}
|
|
original_insert = tree.insert
|
|
|
|
def wrapped_insert(params):
|
|
captured["prev_prefix_len"] = params.prev_prefix_len
|
|
captured["is_bigram"] = params.key.is_bigram
|
|
captured["key_len"] = len(params.key)
|
|
return original_insert(params)
|
|
|
|
tree.insert = wrapped_insert
|
|
tree.cache_finished_req(req, is_insert=True)
|
|
|
|
self.assertEqual(captured["prev_prefix_len"], req.cache_protected_len)
|
|
self.assertTrue(captured["is_bigram"])
|
|
self.assertEqual(captured["key_len"], len(req.origin_input_ids) - 1)
|
|
|
|
# Case 2: is_insert=False should free [cache_protected_len:page_aligned_len]
|
|
# even when len(prefix_indices) is intentionally larger.
|
|
req2 = _DummyReq()
|
|
req2.req_pool_idx = 1
|
|
req2.origin_input_ids = array("q", [11, 12, 13, 14, 15, 16])
|
|
req2.output_ids = array("q")
|
|
req2._kv_committed_len = len(req2.origin_input_ids)
|
|
kv_indices2 = allocator.alloc(req2._kv_committed_len)
|
|
req_to_token_pool.write(
|
|
(req2.req_pool_idx, slice(0, req2._kv_committed_len)), kv_indices2
|
|
)
|
|
req2.extra_key = None
|
|
req2.last_node = tree.root_node
|
|
req2.swa_uuid_for_lock = None
|
|
req2.swa_evicted_seqlen = 0
|
|
req2.cache_protected_len = 1
|
|
req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device)
|
|
|
|
freed_lens = []
|
|
original_free = allocator.free
|
|
|
|
def wrapped_free(indices):
|
|
freed_lens.append(int(indices.numel()))
|
|
return original_free(indices)
|
|
|
|
allocator.free = wrapped_free
|
|
tree.cache_finished_req(req2, is_insert=False)
|
|
|
|
# EAGLE + page_size=1 => page_aligned_len = committed_len - 1 = 5
|
|
# Expected frees:
|
|
# overlap range [1:5] -> 4
|
|
# tail range [5:] -> 1
|
|
self.assertEqual(freed_lens, [4, 1])
|
|
|
|
|
|
# Optimization: SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.
|
|
# Splits a freshly-inserted leaf at the (page-aligned) sliding-window
|
|
# boundary so a future inc_lock_ref protects only ~sliding_window_size SWA
|
|
# tokens instead of the whole chunked-prefill chain.
|
|
class TestSWASplitLeafOnInsert(CustomTestCase):
|
|
def _insert_and_lock(self, *, window, page_size, leaf_len, flag_on):
|
|
tree, allocator, _ = _build_swa_tree(
|
|
is_eagle=False,
|
|
kv_size=128,
|
|
kv_size_swa=64,
|
|
sliding_window_size=window,
|
|
page_size=page_size,
|
|
)
|
|
token_ids = list(range(leaf_len))
|
|
with envs.SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.override(flag_on):
|
|
leaf = _insert_chain(tree, allocator, token_ids)
|
|
result = tree.inc_lock_ref(leaf)
|
|
return tree, leaf, result
|
|
|
|
def test_flag_off_protects_full_leaf(self):
|
|
tree, leaf, _ = self._insert_and_lock(
|
|
window=4, page_size=1, leaf_len=12, flag_on=False
|
|
)
|
|
self.assertEqual(len(leaf.value), 12)
|
|
self.assertEqual(tree.swa_protected_size_, 12)
|
|
|
|
def test_flag_on_caps_protection_at_window(self):
|
|
# (window, page_size, leaf_len, expected_tail_size); leaf_len picked
|
|
# > tail_size and page-aligned for page_size > 1.
|
|
cases = [
|
|
(4, 1, 12, 4),
|
|
(4, 1, 5, 4),
|
|
(1, 1, 5, 1),
|
|
(4, 2, 12, 4),
|
|
(8, 2, 12, 8),
|
|
(4, 4, 12, 4),
|
|
# window NOT page-aligned -> tail rounds up to page boundary.
|
|
(3, 2, 12, 4),
|
|
(5, 4, 12, 8),
|
|
(3, 4, 12, 4),
|
|
]
|
|
for window, page_size, leaf_len, expected_tail in cases:
|
|
with self.subTest(window=window, page_size=page_size, leaf_len=leaf_len):
|
|
self.assertEqual(_expected_tail_size(window, page_size), expected_tail)
|
|
tree, leaf, _ = self._insert_and_lock(
|
|
window=window,
|
|
page_size=page_size,
|
|
leaf_len=leaf_len,
|
|
flag_on=True,
|
|
)
|
|
self.assertEqual(len(leaf.value), expected_tail)
|
|
self.assertEqual(tree.swa_protected_size_, expected_tail)
|
|
|
|
def test_flag_on_no_split_when_leaf_within_window(self):
|
|
# leaf_len <= tail_size: split must no-op.
|
|
cases = [
|
|
(4, 1, 4),
|
|
(4, 1, 3),
|
|
(4, 2, 4),
|
|
(3, 2, 4),
|
|
(8, 2, 4),
|
|
(4, 4, 4),
|
|
]
|
|
for window, page_size, leaf_len in cases:
|
|
with self.subTest(window=window, page_size=page_size, leaf_len=leaf_len):
|
|
tree, leaf, _ = self._insert_and_lock(
|
|
window=window,
|
|
page_size=page_size,
|
|
leaf_len=leaf_len,
|
|
flag_on=True,
|
|
)
|
|
self.assertEqual(len(leaf.value), leaf_len)
|
|
self.assertEqual(tree.swa_protected_size_, leaf_len)
|
|
|
|
def test_match_prefix_returns_full_chain_after_split(self):
|
|
tree, allocator, _ = _build_swa_tree(
|
|
is_eagle=False,
|
|
kv_size=128,
|
|
kv_size_swa=64,
|
|
sliding_window_size=4,
|
|
page_size=1,
|
|
)
|
|
token_ids = list(range(12))
|
|
with envs.SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.override(True):
|
|
inserted_leaf = _insert_chain(tree, allocator, token_ids)
|
|
self.assertEqual(len(inserted_leaf.value), 4)
|
|
match = tree.match_prefix(
|
|
MatchPrefixParams(key=RadixKey(array("q", token_ids)))
|
|
)
|
|
self.assertEqual(match.device_indices.shape[0], 12)
|
|
self.assertIs(match.last_device_node, inserted_leaf)
|
|
|
|
def test_dec_lock_ref_after_split_balances_to_zero(self):
|
|
tree, leaf, result = self._insert_and_lock(
|
|
window=4, page_size=1, leaf_len=12, flag_on=True
|
|
)
|
|
self.assertEqual(tree.swa_protected_size_, 4)
|
|
self.assertEqual(tree.full_protected_size_, 12)
|
|
|
|
tree.dec_lock_ref(
|
|
leaf,
|
|
params=DecLockRefParams(swa_uuid_for_lock=result.swa_uuid_for_lock),
|
|
)
|
|
|
|
self.assertEqual(tree.swa_protected_size_, 0)
|
|
self.assertEqual(tree.full_protected_size_, 0)
|
|
tree.sanity_check()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|