Files
sglang/test/registered/unit/mem_cache/test_radix_key_views.py
leavelet 0d065a8ab0 Speed up Python bigram key conversion with C-level zip
The EAGLE bigram fallback built 64K-token keys with an index-based
Python comprehension (6.3 ms per call at 64K tokens). zip + islice
runs the pairing loop in C and avoids copying the shifted operand:
3.6-4.2 ms per call (~1.6x). Output is byte-identical (same tuples);
the tai-kernel fast path is unaffected.

Packing bigrams into int64s via numpy was measured and rejected: the
list<->array boxing makes it slower (4.1 ms) than zip until token ids
are numpy end-to-end, and it would change HiCache storage hash inputs.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-10 04:56:24 +00:00

234 lines
8.8 KiB
Python

"""
Unit tests for RadixKey zero-copy view semantics.
RadixKey slicing returns O(1) views over a shared backing list instead of
copying. These tests pin down the semantics that the radix tree relies on:
- Slicing/iteration/indexing behave exactly like list slicing.
- Views share the backing list; compacted() detaches them.
- Key matching and child-key derivation are view-position independent.
- Keys stored in tree nodes are always compact (never views), so lock-ref
walks, eviction, splits, and controller-thread reads never observe a key
that pins a transient backing list.
Usage:
python -m pytest test_radix_key_views.py -v
"""
import sys
import unittest.mock
for _mod in ("sgl_kernel", "sgl_kernel.kvcacheio"):
if _mod not in sys.modules:
sys.modules[_mod] = unittest.mock.MagicMock()
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
# CPU-based unit test, runs quickly on any GPU runner
register_cuda_ci(est_time=5, suite="stage-b-test-1-gpu-small")
register_amd_ci(est_time=5, suite="stage-b-test-1-gpu-small-amd")
import random
import unittest
import torch
from sglang.srt.mem_cache.base_prefix_cache import InsertParams, MatchPrefixParams
from sglang.srt.mem_cache.radix_cache import (
RadixCache,
RadixKey,
_key_match_page_size1,
_key_match_paged,
get_child_key,
maybe_bigram_convert,
)
class TestBigramConversionEquivalence(unittest.TestCase):
def test_zip_impl_matches_index_comprehension(self):
from sglang.srt.mem_cache.utils import _python_convert_to_bigram_key
rng = random.Random(7)
for n in (0, 1, 2, 3, 17, 256):
tokens = [rng.randrange(1000) for _ in range(n)]
want = [(tokens[i], tokens[i + 1]) for i in range(max(len(tokens) - 1, 0))]
self.assertEqual(_python_convert_to_bigram_key(tokens), want)
# Already-converted (tuple) input is returned unchanged.
pairs = [(1, 2), (2, 3)]
self.assertEqual(_python_convert_to_bigram_key(pairs), pairs)
class TestRadixKeyViewSemantics(unittest.TestCase):
def test_slice_is_view_sharing_backing(self):
backing = list(range(100))
key = RadixKey(backing)
view = key[10:50]
self.assertIs(view._tokens, key._tokens)
self.assertEqual(view.token_ids, backing[10:50])
self.assertEqual(len(view), 40)
def test_nested_slicing_matches_list_semantics(self):
rng = random.Random(0)
tokens = [rng.randrange(1000) for _ in range(257)]
key = RadixKey(tokens)
ref = list(tokens)
for _ in range(200):
if len(ref) == 0:
break
a = rng.randrange(0, len(ref) + 1)
b = rng.randrange(a, len(ref) + 1)
key = key[a:b]
ref = ref[a:b]
self.assertEqual(key.token_ids, ref)
self.assertEqual(len(key), len(ref))
self.assertEqual(list(key), ref)
def test_open_ended_and_negative_slices(self):
tokens = list(range(20))
key = RadixKey(tokens)
self.assertEqual(key[5:].token_ids, tokens[5:])
self.assertEqual(key[:7].token_ids, tokens[:7])
self.assertEqual(key[-4:].token_ids, tokens[-4:])
self.assertEqual(key[3:1].token_ids, [])
def test_int_indexing(self):
tokens = [7, 8, 9, 10]
key = RadixKey(tokens)[1:]
item = key[0]
self.assertEqual(item.token_ids, [8])
self.assertEqual(key[-1].token_ids, [10])
def test_strided_slicing_rejected(self):
key = RadixKey(list(range(10)))
with self.assertRaises(ValueError):
key[::2]
def test_compacted_detaches_views_and_is_identity_for_full_keys(self):
backing = list(range(50))
key = RadixKey(backing, extra_key="lora1", is_bigram=False)
self.assertIs(key.compacted(), key)
view = key[10:30]
compact = view.compacted()
self.assertIsNot(compact._tokens, backing)
self.assertEqual(compact.token_ids, backing[10:30])
self.assertEqual(compact.extra_key, "lora1")
# Mutating the original backing list must not affect the compacted key.
backing[15] = -1
self.assertNotIn(-1, compact.token_ids)
def test_token_ids_setter_resets_view(self):
key = RadixKey(list(range(30)))[5:25]
key.token_ids = [1, 2, 3]
self.assertEqual(len(key), 3)
self.assertEqual(key.token_ids, [1, 2, 3])
def test_length_frozen_at_construction(self):
backing = [1, 2, 3]
key = RadixKey(backing)
backing.append(4)
self.assertEqual(len(key), 3)
self.assertEqual(key.token_ids, [1, 2, 3])
def test_extra_key_and_bigram_preserved_through_slices(self):
key = RadixKey([(1, 2), (2, 3), (3, 4)], extra_key="salt", is_bigram=True)
view = key[1:]
self.assertEqual(view.extra_key, "salt")
self.assertTrue(view.is_bigram)
self.assertTrue(view.compacted().is_bigram)
def test_bigram_convert_on_view(self):
key = RadixKey(list(range(10)))[2:8]
converted, _ = maybe_bigram_convert(True, key)
self.assertTrue(converted.is_bigram)
self.assertEqual(
converted.token_ids[0][0] if converted.token_ids else None,
2,
)
def test_key_match_is_view_position_independent(self):
rng = random.Random(1)
base = [rng.randrange(50) for _ in range(512)]
other = list(base[:300]) + [999] + base[301:]
for page_size in (1, 16, 64):
for off in (0, 1, 63, 128):
k0 = RadixKey([0] * off + base)[off:]
k1 = RadixKey(other)
k0_ref = RadixKey(list(base))
if page_size == 1:
got = _key_match_page_size1(k0, k1)
want = _key_match_page_size1(k0_ref, k1)
else:
got = _key_match_paged(k0, k1, page_size)
want = _key_match_paged(k0_ref, k1, page_size)
self.assertEqual(got, want, f"{page_size=} {off=}")
def test_get_child_key_on_views(self):
tokens = list(range(200))
key = RadixKey(tokens, extra_key="e")
view = key[64:]
self.assertEqual(get_child_key(view, 64), get_child_key(view.compacted(), 64))
self.assertEqual(get_child_key(view, 64), ("e", tuple(tokens[64:128])))
self.assertEqual(get_child_key(view, 1), ("e", 64))
def test_tree_node_keys_are_always_compact(self):
"""Node-stored keys must never be views: a view would pin the transient
request key's backing list, and lock/update paths (inc_lock_ref,
eviction, splits) would read keys whose backing outlives the match."""
rng = random.Random(2)
cache = RadixCache.create_simulated(page_size=4)
seqs = []
base = [rng.randrange(8) for _ in range(256)]
for _ in range(40):
# Shared prefixes of random length force splits at random points.
cut = rng.randrange(8, 256, 4)
seq = base[:cut] + [rng.randrange(8) for _ in range(rng.randrange(4, 64))]
seqs.append(seq)
cache.insert(
InsertParams(
key=RadixKey(seq),
value=torch.arange(len(seq), dtype=torch.int64),
)
)
cache.match_prefix(MatchPrefixParams(key=RadixKey(list(seq))))
stack = [cache.root_node]
checked = 0
while stack:
node = stack.pop()
stack.extend(node.children.values())
if node is cache.root_node or node.key is None:
continue
checked += 1
self.assertEqual(node.key._start, 0, f"node {node.id} key is a view")
self.assertEqual(
node.key._stop,
len(node.key._tokens),
f"node {node.id} key is a view",
)
self.assertGreater(checked, 10)
def test_match_and_insert_equivalence_random(self):
"""Differential test: tree behavior identical for fresh keys vs views."""
rng = random.Random(3)
cache = RadixCache.create_simulated(page_size=2)
for _ in range(50):
seq = [rng.randrange(4) for _ in range(rng.randrange(2, 40, 2))]
cache.insert(
InsertParams(
key=RadixKey(seq),
value=torch.arange(len(seq), dtype=torch.int64),
)
)
probe = [rng.randrange(4) for _ in range(rng.randrange(2, 40, 2))]
res_fresh = cache.match_prefix(MatchPrefixParams(key=RadixKey(list(probe))))
padded = RadixKey([9] * 6 + list(probe))[6:]
res_view = cache.match_prefix(MatchPrefixParams(key=padded))
self.assertTrue(
torch.equal(res_fresh.device_indices, res_view.device_indices)
)
if __name__ == "__main__":
unittest.main()