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>
234 lines
8.8 KiB
Python
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()
|