Make RadixKey slicing zero-copy to fix quadratic match_prefix

RadixKey.__getitem__ copied the token list on every slice, and the
match/insert tree walks re-slice the remaining key at every node hop,
making a single match O(len * hops) — quadratic for long cached
prefixes. Slices now return O(1) offset-based views over a shared
backing list; the key-match and child-key functions index the backing
list directly so views are never materialized on the hot path.

Keys stored in tree nodes are compacted at every store site (same cost
as the old copying slices), so lock-ref walks, eviction, splits, and
controller-thread reads never observe a key pinning a transient
backing list. Node-key compactness is enforced by a tree-walk test.

Microbenchmark (64K-token full hit, page_size=64):
  16-node path: 1.86 -> 0.87 ms/match (2.1x)
  128-node path: 9.80 -> 1.19 ms/match (8.2x), insert re-walk 6.9x

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-06-10 04:28:03 +00:00
co-authored by Claude Fable 5
parent ffff715f00
commit 8979c81a22
7 changed files with 311 additions and 31 deletions
@@ -0,0 +1,219 @@
"""
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 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()