项目文件夹

文件
wehub-resource-sync 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
chore: import upstream snapshot with attribution
2026-07-13 12:38:16 +08:00

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()