[RadixTree][1/N Refactor]: Support unified match_prefix params (#17142)

Co-authored-by: yizhang2077 <1109276519@qq.com>
Co-authored-by: pansicheng <sicheng.pan.chn@gmail.com>
This commit is contained in:
zhangheng
2026-01-19 22:39:40 +08:00
committed by GitHub
co-authored by yizhang2077 pansicheng
parent ce8a6ac690
commit 20b0523eca
13 changed files with 117 additions and 64 deletions
@@ -31,6 +31,7 @@ import unittest.mock
import torch
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
# Test constants
@@ -294,12 +295,12 @@ class TestRadixCache(unittest.TestCase):
self.assertEqual(cache.evictable_size(), 3)
# Test match_prefix
result = cache.match_prefix(RadixKey([1, 2, 3]))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
self.assertEqual(len(result.device_indices), 3)
torch.testing.assert_close(result.device_indices, value)
# Test partial match
result = cache.match_prefix(RadixKey([1, 2]))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2])))
self.assertEqual(len(result.device_indices), 2)
torch.testing.assert_close(
result.device_indices, torch.tensor([10, 20], dtype=torch.int64)
@@ -402,10 +403,12 @@ class TestRadixCache(unittest.TestCase):
)
# Keys with different extra_key should not match each other
result1 = cache.match_prefix(RadixKey([1, 2, 3], "key1"))
result2 = cache.match_prefix(RadixKey([1, 2, 3], "key2"))
result3 = cache.match_prefix(RadixKey([1, 2, 3], None))
result4 = cache.match_prefix(RadixKey([1, 2, 3], "nonexistent"))
result1 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], "key1")))
result2 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], "key2")))
result3 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], None)))
result4 = cache.match_prefix(
MatchPrefixParams(key=RadixKey([1, 2, 3], "nonexistent"))
)
# Each should match only its own data
self.assertEqual(len(result1.device_indices), 3)
@@ -434,7 +437,7 @@ class TestRadixCache(unittest.TestCase):
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
# Get node
result = cache.match_prefix(RadixKey([1, 2, 3]))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
node = result.last_device_node
initial_evictable = cache.evictable_size()
@@ -485,7 +488,7 @@ class TestRadixCache(unittest.TestCase):
tokens = list(range(sequence_length))
cache.insert(RadixKey(tokens), torch.tensor(tokens, dtype=torch.int64))
result = cache.match_prefix(RadixKey(tokens))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertGreater(len(result.device_indices), 0)
# Match length should be page-aligned
@@ -541,23 +544,25 @@ class TestRadixCache(unittest.TestCase):
# Match that causes a split inside an existing node:
# take first 4 tokens of seq1, then diverge.
query1 = [1, 2, 3, 4, 999, 1000]
result1 = cache.match_prefix(RadixKey(query1))
result1 = cache.match_prefix(MatchPrefixParams(key=RadixKey(query1)))
torch.testing.assert_close(result1.device_indices, val1[:4])
# No data change after structural split during matching.
self.assertEqual(cache.total_size(), baseline_total)
# Full match of the long sequence still returns the full indices.
result_full = cache.match_prefix(RadixKey(seq1))
result_full = cache.match_prefix(MatchPrefixParams(key=RadixKey(seq1)))
torch.testing.assert_close(result_full.device_indices, val1)
# Another split deeper on the path (after matching 6 tokens, then diverge).
query2 = [1, 2, 3, 4, 5, 6, 777, 888]
result2 = cache.match_prefix(RadixKey(query2))
result2 = cache.match_prefix(MatchPrefixParams(key=RadixKey(query2)))
torch.testing.assert_close(result2.device_indices, val1[:6])
self.assertEqual(cache.total_size(), baseline_total)
# Matching the short diverging branch should return exactly its indices.
result_branch = cache.match_prefix(RadixKey(seq2))
result_branch = cache.match_prefix(
MatchPrefixParams(key=RadixKey(seq2))
)
torch.testing.assert_close(result_branch.device_indices, val2)
def test_hash_value_storage(self):