[Refactor] Replace page_align_keys helper with RadixKey.page_aligned method (#23107)

This commit is contained in:
Liangsheng Yin
2026-04-20 18:10:42 -07:00
committed by GitHub
parent 712b01d875
commit c7a4ebf3c8
10 changed files with 126 additions and 135 deletions
@@ -156,10 +156,11 @@ class TestMamba(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(req1_token_ids)
result = tree.insert(
InsertParams(
key=RadixKey(req1_token_ids),
value=req1_kv_indices,
key=key,
value=req1_kv_indices[: len(key)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
@@ -173,10 +174,11 @@ class TestMamba(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(req2_token_ids)
result = tree.insert(
InsertParams(
key=RadixKey(req2_token_ids),
value=req2_kv_indices,
key=key,
value=req2_kv_indices[: len(key)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
@@ -191,10 +193,11 @@ class TestMamba(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(req3_token_ids)
result = tree.insert(
InsertParams(
key=RadixKey(req3_token_ids),
value=req3_kv_indices,
key=key,
value=req3_kv_indices[: len(key)],
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
)
)
@@ -208,10 +211,11 @@ class TestMamba(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(req4_token_ids)
result = tree.insert(
InsertParams(
key=RadixKey(req4_token_ids),
value=req4_kv_indices,
key=key,
value=req4_kv_indices[: len(key)],
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
)
)
@@ -400,10 +404,11 @@ class TestMamba(unittest.TestCase):
# Step 1: Insert [1,2,3] to create first node
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3])
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=allocator.alloc(3),
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
@@ -412,10 +417,11 @@ class TestMamba(unittest.TestCase):
# Step 2: Insert [1,2,3,4,5,6,7] with prev_prefix_len=0 (free all matched)
# Creates tree: [1,2,3] -> [4,5,6,7]
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 3, 4, 5, 6, 7])
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7]),
value=allocator.alloc(7),
key=key2,
value=allocator.alloc(7)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=0,
)
@@ -429,10 +435,11 @@ class TestMamba(unittest.TestCase):
# Matched prefix = 7 (across two nodes: [1,2,3] len=3, [4,5,6,7] len=4)
# Protected [0..1], freed [2..6] = 5 slots, new [7] = 1 slot stored
req3 = make_dummy_req()
key3 = RadixKey([1, 2, 3, 4, 5, 6, 7, 8])
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7, 8]),
value=allocator.alloc(8),
key=key3,
value=allocator.alloc(8)[: len(key3)],
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=2,
)
@@ -445,10 +452,11 @@ class TestMamba(unittest.TestCase):
# Step 4: Insert [1,2,3,4,5,6,7,8,9] with prev_prefix_len=8 (covers all matched)
# Matched prefix = 8, prev_prefix_len=8 => nothing freed
req4 = make_dummy_req()
key4 = RadixKey([1, 2, 3, 4, 5, 6, 7, 8, 9])
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7, 8, 9]),
value=allocator.alloc(9),
key=key4,
value=allocator.alloc(9)[: len(key4)],
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=8,
)
@@ -62,9 +62,7 @@ class TestSLRUAccuracy(unittest.TestCase):
"""Test that SLRU eviction mechanism works correctly"""
# Insert one key-value three times (high frequency access)
frequent_key = RadixKey(
token_ids=[1, 2], extra_key=None
) # High hit rate, should be retained
frequent_key = RadixKey([1, 2]) # High hit rate, should be retained
frequent_val = torch.tensor([10, 20], dtype=torch.int64)
# Insert the frequent key multiple times to increase its hit count
@@ -72,9 +70,7 @@ class TestSLRUAccuracy(unittest.TestCase):
self.cache.insert(InsertParams(key=frequent_key, value=frequent_val))
# Insert first low-frequency key-value pair that should be evicted
first_low_freq_key = RadixKey(
token_ids=[5, 6], extra_key=None
) # Low hit rate, should be evicted
first_low_freq_key = RadixKey([5, 6]) # Low hit rate, should be evicted
first_low_freq_val = torch.tensor([50, 60], dtype=torch.int64)
self.cache.insert(
@@ -84,18 +80,14 @@ class TestSLRUAccuracy(unittest.TestCase):
# Insert other key-values once each (low frequency access) - fill up the cache
other_keys = []
for i in range(4): # Reduce the number to fit in our smaller cache
key = RadixKey(
token_ids=[i + 10], extra_key=None
) # Unique keys for low-frequency items
key = RadixKey([i + 10]) # Unique keys for low-frequency items
val = torch.tensor([i + 100], dtype=torch.int64)
self.cache.insert(InsertParams(key=key, value=val))
other_keys.append(key)
# Now insert more items to trigger evictions
for i in range(6, 10): # Add more items to definitely exceed capacity
key = RadixKey(
token_ids=[i * 2], extra_key=None
) # Different pattern to avoid conflicts
key = RadixKey([i * 2]) # Different pattern to avoid conflicts
val = torch.tensor([i * 200], dtype=torch.int64)
self.cache.insert(InsertParams(key=key, value=val))
@@ -547,10 +547,11 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated(page_size=page_size)
tokens = list(range(sequence_length))
key = RadixKey(tokens)
cache.insert(
InsertParams(
key=RadixKey(tokens),
value=torch.tensor(tokens, dtype=torch.int64),
key=key,
value=torch.tensor(tokens, dtype=torch.int64)[: len(key)],
)
)
@@ -215,9 +215,8 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req1_token_ids), value=req1_kv_indices)
)
key = RadixKey(req1_token_ids)
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
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()}"
@@ -227,9 +226,8 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req2_token_ids), value=req2_kv_indices)
)
key = RadixKey(req2_token_ids)
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
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()}"
@@ -239,9 +237,8 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req3_token_ids), value=req3_kv_indices)
)
key = RadixKey(req3_token_ids)
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
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()}"
@@ -251,9 +248,8 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req4_token_ids), value=req4_kv_indices)
)
key = RadixKey(req4_token_ids)
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
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()}"
@@ -374,9 +370,8 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req1_token_ids), value=req1_kv_indices)
)
key = RadixKey(req1_token_ids)
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 0)
print(
@@ -387,9 +382,8 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req2_token_ids), value=req2_kv_indices)
)
key = RadixKey(req2_token_ids)
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 2)
print(
@@ -400,9 +394,8 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req3_token_ids), value=req3_kv_indices)
)
key = RadixKey(req3_token_ids)
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 0)
print(
@@ -413,9 +406,8 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
result = tree.insert(
InsertParams(key=RadixKey(req4_token_ids), value=req4_kv_indices)
)
key = RadixKey(req4_token_ids)
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 4)
print(
@@ -335,7 +335,8 @@ def _insert_seq(env, seq):
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
env.tree.insert(InsertParams(key=RadixKey(seq), value=v, mamba_value=mamba_val))
key = RadixKey(seq)
env.tree.insert(InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val))
return True
@@ -356,7 +357,10 @@ def _fill_no_evict(env):
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
env.tree.insert(InsertParams(key=RadixKey(seq), value=v, mamba_value=mamba_val))
key = RadixKey(seq)
env.tree.insert(
InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val)
)
inserted += 1
return inserted
@@ -501,8 +505,9 @@ def bench_match_prefix(
queries.append([rng.randint(1, 32000)] * rng.randint(50, 300))
def verify_fn(q):
r1 = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q)))
r2 = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q)))
k = RadixKey(q)
r1 = env.tree.match_prefix(MatchPrefixParams(key=k))
r2 = env.tree.match_prefix(MatchPrefixParams(key=k))
assert len(r1.device_indices) == len(r2.device_indices), "match not idempotent"
warmup = min(20, len(queries) // 10)
@@ -258,10 +258,9 @@ class UnifiedRadixCacheSuite:
def _insert(self, tree, allocator, req_to_token_pool, tokens):
"""Insert tokens, attaching mamba data when the config has mamba."""
params = InsertParams(
key=RadixKey(tokens),
value=self._alloc(allocator, len(tokens)),
)
key = RadixKey(tokens)
value = self._alloc(allocator, len(tokens))
params = InsertParams(key=key, value=value[: len(key)])
if self.cfg.has_mamba:
req = self._make_req(req_to_token_pool)
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
@@ -400,9 +399,11 @@ class UnifiedRadixCacheSuite:
self.assertEqual(allocator.available_size(), initial_avail - len(seq_1p))
# Step 2: insert 2 pages with prev_prefix_len=0 → frees overlap of 1 page
key_2p = RadixKey(seq_2p)
value_2p = self._alloc(allocator, len(seq_2p))
params = InsertParams(
key=RadixKey(seq_2p),
value=self._alloc(allocator, len(seq_2p)),
key=key_2p,
value=value_2p[: len(key_2p)],
prev_prefix_len=0,
)
if self.cfg.has_mamba:
@@ -417,9 +418,11 @@ class UnifiedRadixCacheSuite:
# Step 3: insert 3 pages with prev_prefix_len=len(seq_2p) → nothing freed
avail_before = allocator.available_size()
key_3p = RadixKey(seq_3p)
value_3p = self._alloc(allocator, len(seq_3p))
params = InsertParams(
key=RadixKey(seq_3p),
value=self._alloc(allocator, len(seq_3p)),
key=key_3p,
value=value_3p[: len(key_3p)],
prev_prefix_len=len(seq_2p),
)
if self.cfg.has_mamba:
@@ -582,6 +585,7 @@ class UnifiedRadixCacheSuite:
self.assertIsInstance(child_key, tuple)
def test_paged_match_truncates_unaligned_key(self):
"""match_prefix internally aligns keys to page boundary."""
if self.cfg.page_size == 1:
self.skipTest("page_size > 1 only")
ps = self.cfg.page_size
@@ -589,10 +593,12 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
# Tree truncates unaligned tail internally, so it matches the seq prefix.
unaligned = seq + list(range(9000, 9000 + ps - 1))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(unaligned)))
self.assertEqual(len(m.device_indices), len(seq))
# Below-page-size key aligns to 0 -> no match.
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq[: ps - 1])))
self.assertEqual(len(m.device_indices), 0)