[RadixTree][3/N Refactor]:Support unified insert/evict params (#17401)

This commit is contained in:
zhangheng
2026-01-22 17:36:31 +08:00
committed by GitHub
parent 2262c5c9b5
commit f33022d039
13 changed files with 462 additions and 138 deletions
@@ -6,7 +6,11 @@ import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams
from sglang.srt.mem_cache.base_prefix_cache import (
EvictParams,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
@@ -234,9 +238,14 @@ class TestMamba(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
prefix_len = tree.insert(
RadixKey(req1_token_ids), req1_kv_indices, req1.mamba_pool_idx.unsqueeze(0)
result = tree.insert(
InsertParams(
key=RadixKey(req1_token_ids),
value=req1_kv_indices,
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
)
@@ -246,9 +255,14 @@ class TestMamba(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
prefix_len = tree.insert(
RadixKey(req2_token_ids), req2_kv_indices, req2.mamba_pool_idx.unsqueeze(0)
result = tree.insert(
InsertParams(
key=RadixKey(req2_token_ids),
value=req2_kv_indices,
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
)
@@ -259,9 +273,14 @@ class TestMamba(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
prefix_len = tree.insert(
RadixKey(req3_token_ids), req3_kv_indices, req3.mamba_pool_idx.unsqueeze(0)
result = tree.insert(
InsertParams(
key=RadixKey(req3_token_ids),
value=req3_kv_indices,
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
)
@@ -271,9 +290,14 @@ class TestMamba(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
prefix_len = tree.insert(
RadixKey(req4_token_ids), req4_kv_indices, req4.mamba_pool_idx.unsqueeze(0)
result = tree.insert(
InsertParams(
key=RadixKey(req4_token_ids),
value=req4_kv_indices,
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
)
@@ -281,12 +305,18 @@ class TestMamba(unittest.TestCase):
tree.pretty_print()
full_num_tokens = 1
print(f"evicting {full_num_tokens} full token")
tree.evict(full_num_tokens=full_num_tokens)
result = tree.evict(EvictParams(num_tokens=full_num_tokens))
assert (
result.num_tokens_evicted >= full_num_tokens
), f"evicted {result.num_tokens_evicted} full tokens, expected {full_num_tokens}"
tree.pretty_print()
mamba_num = 1
print(f"evicting {mamba_num} mamba")
tree.evict_mamba(mamba_num=mamba_num)
result = tree.evict(EvictParams(num_tokens=0, mamba_num=mamba_num))
assert (
result.mamba_num_evicted >= mamba_num
), f"evicted {result.mamba_num_evicted} mamba states, expected {mamba_num}"
tree.pretty_print()
req5_token_ids = [1, 2, 3, 4, 5]
@@ -317,7 +347,10 @@ class TestMamba(unittest.TestCase):
mamba_num = 1
print(f"evicting {mamba_num} mamba")
tree.evict_mamba(mamba_num=mamba_num)
result = tree.evict(EvictParams(num_tokens=0, mamba_num=mamba_num))
assert (
result.mamba_num_evicted >= mamba_num
), f"evicted {result.mamba_num_evicted} mamba states, expected {mamba_num}"
tree.pretty_print()
req8_token_ids = [1, 2, 3, 4, 5, 60, 70]
@@ -31,7 +31,12 @@ 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.base_prefix_cache import (
EvictParams,
EvictResult,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
# Test constants
@@ -266,7 +271,12 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated()
# Insert some data
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
self.assertGreater(cache.total_size(), 0)
# Reset
@@ -283,7 +293,8 @@ class TestRadixCache(unittest.TestCase):
key = RadixKey([1, 2, 3])
value = torch.tensor([10, 20, 30], dtype=torch.int64)
prefix_len = cache.insert(key, value)
result = cache.insert(InsertParams(key=key, value=value))
prefix_len = result.prefix_len
if disable_cache:
self.assertEqual(prefix_len, 0)
@@ -311,7 +322,8 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated()
key = RadixKey([1, 2, 3])
prefix_len = cache.insert(key, None)
result = cache.insert(InsertParams(key=key, value=None))
prefix_len = result.prefix_len
# When None is passed, it should create value from token_ids
self.assertEqual(prefix_len, 0)
@@ -323,10 +335,19 @@ class TestRadixCache(unittest.TestCase):
self.assertEqual(cache.total_size(), 0)
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
self.assertEqual(cache.total_size(), 3)
cache.insert(RadixKey([4, 5]), torch.tensor([40, 50], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([4, 5]), value=torch.tensor([40, 50], dtype=torch.int64)
)
)
self.assertEqual(cache.total_size(), 5)
def test_kv_cache_events(self):
@@ -344,7 +365,7 @@ class TestRadixCache(unittest.TestCase):
)
# Insert data
cache.insert(RadixKey([1, 2, 3, 4, 5]), None)
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 5]), value=None))
# Take events
events = cache.take_events()
@@ -371,8 +392,19 @@ class TestRadixCache(unittest.TestCase):
)
# Insert and then evict data
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
cache.evict(3)
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
result = cache.evict(EvictParams(num_tokens=3))
self.assertIsInstance(result, EvictResult)
self.assertGreaterEqual(
result.num_tokens_evicted,
3,
f"evicted {result.num_tokens_evicted} tokens, expected at least 3",
)
# Take events - should include both store and remove events
events = cache.take_events()
@@ -393,13 +425,22 @@ class TestRadixCache(unittest.TestCase):
# Insert same token sequence with different extra keys
cache.insert(
RadixKey([1, 2, 3], "key1"), torch.tensor([10, 20, 30], dtype=torch.int64)
InsertParams(
key=RadixKey([1, 2, 3], "key1"),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
cache.insert(
RadixKey([1, 2, 3], "key2"), torch.tensor([40, 50, 60], dtype=torch.int64)
InsertParams(
key=RadixKey([1, 2, 3], "key2"),
value=torch.tensor([40, 50, 60], dtype=torch.int64),
)
)
cache.insert(
RadixKey([1, 2, 3], None), torch.tensor([70, 80, 90], dtype=torch.int64)
InsertParams(
key=RadixKey([1, 2, 3], None),
value=torch.tensor([70, 80, 90], dtype=torch.int64),
)
)
# Keys with different extra_key should not match each other
@@ -434,7 +475,12 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated()
# Insert sequence
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
# Get node
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
@@ -461,13 +507,27 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated(mock_allocator=mock_allocator)
# Insert sequences
cache.insert(RadixKey([1, 2]), torch.tensor([10, 20], dtype=torch.int64))
cache.insert(RadixKey([3, 4]), torch.tensor([30, 40], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2]), value=torch.tensor([10, 20], dtype=torch.int64)
)
)
cache.insert(
InsertParams(
key=RadixKey([3, 4]), value=torch.tensor([30, 40], dtype=torch.int64)
)
)
initial_size = cache.total_size()
# Evict some tokens
cache.evict(2)
result = cache.evict(EvictParams(num_tokens=2))
self.assertIsInstance(result, EvictResult)
self.assertGreaterEqual(
result.num_tokens_evicted,
2,
f"evicted {result.num_tokens_evicted} tokens, expected at least 2",
)
# Should have called free and reduced size
mock_allocator.free.assert_called()
@@ -486,7 +546,12 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated(page_size=page_size)
tokens = list(range(sequence_length))
cache.insert(RadixKey(tokens), torch.tensor(tokens, dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey(tokens),
value=torch.tensor(tokens, dtype=torch.int64),
)
)
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
self.assertGreater(len(result.device_indices), 0)
@@ -499,7 +564,12 @@ class TestRadixCache(unittest.TestCase):
"""Test pretty_print produces output."""
cache = RadixCache.create_simulated()
cache.insert(RadixKey([1, 2, 3]), torch.tensor([10, 20, 30], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
# Just test that it doesn't crash
try:
@@ -511,8 +581,16 @@ class TestRadixCache(unittest.TestCase):
"""Test all_values_flatten method."""
cache = RadixCache.create_simulated()
cache.insert(RadixKey([1, 2]), torch.tensor([10, 20], dtype=torch.int64))
cache.insert(RadixKey([3, 4]), torch.tensor([30, 40], dtype=torch.int64))
cache.insert(
InsertParams(
key=RadixKey([1, 2]), value=torch.tensor([10, 20], dtype=torch.int64)
)
)
cache.insert(
InsertParams(
key=RadixKey([3, 4]), value=torch.tensor([30, 40], dtype=torch.int64)
)
)
all_values = cache.all_values_flatten()
self.assertEqual(len(all_values), 4)
@@ -529,12 +607,12 @@ class TestRadixCache(unittest.TestCase):
# Insert a long sequence that will be split later.
seq1 = [1, 2, 3, 4, 5, 6, 7, 8]
val1 = torch.tensor([x * 10 for x in seq1], dtype=torch.int64)
cache.insert(RadixKey(seq1), val1)
cache.insert(InsertParams(key=RadixKey(seq1), value=val1))
# Insert a diverging branch to create an internal node on the path.
seq2 = [1, 2, 9, 10]
val2 = torch.tensor([x * 10 for x in seq2], dtype=torch.int64)
cache.insert(RadixKey(seq2), val2)
cache.insert(InsertParams(key=RadixKey(seq2), value=val2))
print(cache.pretty_print())
baseline_total = cache.total_size()
@@ -573,7 +651,7 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence
cache.insert(RadixKey([1, 2, 3, 4, 5, 6, 7, 8]), None)
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 5, 6, 7, 8]), value=None))
# Trigger event emission to compute hash_value lazily
cache.take_events()
@@ -599,7 +677,7 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence with repeating token pattern: [1,2,3,4, 1,2,3,4]
cache.insert(RadixKey([1, 2, 3, 4, 1, 2, 3, 4]), None)
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 1, 2, 3, 4]), value=None))
events = cache.take_events()
block_stored_events = [e for e in events if isinstance(e, BlockStored)]
@@ -633,11 +711,11 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence that will cause a split
cache.insert(RadixKey([1, 2, 3, 4]), None)
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4]), value=None))
cache.take_events() # Clear events and compute hash_value for first node
# Insert a diverging sequence that will cause a split at page boundary
cache.insert(RadixKey([1, 2, 5, 6]), None)
cache.insert(InsertParams(key=RadixKey([1, 2, 5, 6]), value=None))
cache.take_events() # Trigger event emission to compute hash_value
# Find the split node
@@ -674,7 +752,7 @@ class TestRadixCache(unittest.TestCase):
cache: RadixCache = RadixCache.create_simulated()
for key, value in zip(keys, values):
cache.insert(RadixKey(key), value)
cache.insert(InsertParams(key=RadixKey(key), value=value))
del values
@@ -2,7 +2,12 @@ import unittest
import torch
from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams
from sglang.srt.mem_cache.base_prefix_cache import (
EvictParams,
EvictResult,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
@@ -140,7 +145,10 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req1_token_ids), req1_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req1_token_ids), value=req1_kv_indices)
)
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()}"
)
@@ -149,7 +157,10 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req2_token_ids), req2_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req2_token_ids), value=req2_kv_indices)
)
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()}"
)
@@ -158,7 +169,10 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req3_token_ids), req3_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req3_token_ids), value=req3_kv_indices)
)
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()}"
)
@@ -167,7 +181,10 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req4_token_ids), req4_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req4_token_ids), value=req4_kv_indices)
)
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()}"
)
@@ -175,17 +192,23 @@ class TestSWA(unittest.TestCase):
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(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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]
@@ -277,7 +300,10 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req1_token_ids), req1_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req1_token_ids), value=req1_kv_indices)
)
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()}"
@@ -287,7 +313,10 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req2_token_ids), req2_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req2_token_ids), value=req2_kv_indices)
)
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()}"
@@ -297,7 +326,10 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req3_token_ids), req3_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req3_token_ids), value=req3_kv_indices)
)
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()}"
@@ -307,7 +339,10 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
prefix_len = tree.insert(RadixKey(req4_token_ids), req4_kv_indices)
result = tree.insert(
InsertParams(key=RadixKey(req4_token_ids), value=req4_kv_indices)
)
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()}"
@@ -316,17 +351,41 @@ class TestSWA(unittest.TestCase):
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(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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")
tree.evict(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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")
tree.evict(full_num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
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]