[RaidxTree Refactor]: Support Unified HybridRadixTree V2 (#21206)

Co-authored-by: ispobock <ispobaoke@gmail.com>
Co-authored-by: pansicheng <sicheng.pan.chn@gmail.com>
Co-authored-by: yizhang2077 <1109276519@qq.com>
Co-authored-by: xiezhq-hermann <xiezhq@stanford.edu>
This commit is contained in:
Zhangheng
2026-04-13 10:28:22 +08:00
committed by GitHub
co-authored by ispobock pansicheng yizhang2077 xiezhq-hermann
parent 5593539942
commit bc59cc0f96
15 changed files with 4707 additions and 1 deletions
@@ -0,0 +1,266 @@
import random
import unittest
from types import SimpleNamespace
from urllib.parse import urlparse
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kl_multiturn_utils import (
get_input_ids,
make_mamba_decode_assert,
make_mamba_prefill_assert,
test_input_output_logprobs_match_decode_cache_hit_helper,
test_input_output_logprobs_match_helper,
test_input_output_logprobs_match_prefill_cache_hit_helper,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
def _random_suffixes(n, length, seed):
"""Generate n random token-id lists of the given length."""
rng = random.Random(seed)
return [[rng.randint(1, 30000) for _ in range(length)] for _ in range(n)]
MAMBA_MODEL = "Qwen/Qwen3-Next-80B-A3B-Instruct"
MAMBA_CHUNK_SIZE = 64
MAMBA_TRACK_INTERVAL = 128
SWA_MODEL = "openai/gpt-oss-20b"
FULL_MODEL = "Qwen/Qwen3-32B"
register_cuda_ci(est_time=1200, suite="stage-c-test-4-gpu-h100")
class UnifiedRadixTreeTestMixin:
"""Mixin: gsm8k、mmlu and multi-turn KL tests with multi-branch interleaving."""
kl_threshold: float = 0.003
max_new_tokens: int = 512
num_groups: int = 3
branches_per_group: int = 3
prefix_len: int = 512
prefill_cache_assert = None
decode_cache_assert = None
gsm8k_threshold: float = 0.93
mmlu_threshold: float = 0.8
num_gsm8k_questions: int = 200
def test_gsm8k(self):
"""Few-shot GSM8K math reasoning accuracy."""
from sglang.test.few_shot_gsm8k import run_eval as run_few_shot_gsm8k
url = urlparse(self.base_url)
args = SimpleNamespace(
num_shots=10,
data_path=None,
num_questions=self.num_gsm8k_questions,
max_new_tokens=16000,
parallel=128,
host=f"http://{url.hostname}",
port=int(url.port),
)
metrics = run_few_shot_gsm8k(args)
print(
f"[{self.__class__.__name__}] GSM8K accuracy: {metrics['accuracy']:.3f} "
f"(threshold: {self.gsm8k_threshold})"
)
self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold)
def test_mmlu(self):
"""Simple-evals MMLU multi-task accuracy."""
from sglang.test.run_eval import run_eval as run_simple_eval
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="mmlu",
num_examples=64,
num_threads=32,
)
metrics = run_simple_eval(args)
print(
f"[{self.__class__.__name__}] MMLU score: {metrics['score']:.3f} "
f"(threshold: {self.mmlu_threshold})"
)
self.assertGreaterEqual(metrics["score"], self.mmlu_threshold)
def test_multiturn_logprobs_match(self):
"""Helper 1: 3-turn, no explicit cache seeding."""
ids = self.input_ids[:4]
n = len(ids)
t2 = _random_suffixes(n, 512, seed=100)
t3 = _random_suffixes(n, 256, seed=200)
test_input_output_logprobs_match_helper(
self.base_url,
self.model,
self.kl_threshold,
ids,
turn_suffixes=[t2, t3],
assert_decode_cached_tokens=self.decode_cache_assert,
max_new_tokens=self.max_new_tokens,
)
def test_multiturn_prefill_cache_hit_branching(self):
"""Helper 2: prefill hit + 2 decode-hit turns, multi-branch interleaved."""
num_groups = self.num_groups
branches = self.branches_per_group
n = num_groups * branches
rng = random.Random(456)
prefix_ids, full_ids = [], []
for g in range(num_groups):
prefix = self.input_ids[g][: self.prefix_len]
for b in range(branches):
suffix = [rng.randint(1, 30000) for _ in range(256 + b * 64)]
prefix_ids.append(list(prefix))
full_ids.append(prefix + suffix)
t2 = _random_suffixes(n, 512, seed=789)
t3 = _random_suffixes(n, 256, seed=890)
test_input_output_logprobs_match_prefill_cache_hit_helper(
self.base_url,
self.model,
self.kl_threshold,
prefix_input_ids=prefix_ids,
full_input_ids=full_ids,
turn_suffixes=[t2, t3],
assert_prefill_cached_tokens=self.prefill_cache_assert,
assert_decode_cached_tokens=self.decode_cache_assert,
branches_per_group=branches,
max_new_tokens=self.max_new_tokens,
)
def test_multiturn_decode_cache_hit_branching(self):
"""Helper 3: 3-turn decode hit, multi-branch interleaved."""
num_groups = self.num_groups
branches = self.branches_per_group
n = num_groups * branches
first_turn = []
for g in range(num_groups):
base = self.input_ids[g][: self.prefix_len]
for _ in range(branches):
first_turn.append(list(base))
t2 = _random_suffixes(n, 512, seed=300)
t3 = _random_suffixes(n, 256, seed=400)
test_input_output_logprobs_match_decode_cache_hit_helper(
self.base_url,
self.model,
self.kl_threshold,
first_turn,
turn_suffixes=[t2, t3],
assert_decode_cached_tokens=self.decode_cache_assert,
branches_per_group=branches,
max_new_tokens=self.max_new_tokens,
)
class TestUnifiedFullRadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""Full attention."""
kl_threshold = 0.0025
@classmethod
def setUpClass(cls):
cls.model = FULL_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--mem-fraction-static",
"0.80",
"--page-size",
"64",
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
class TestUnifiedMambaRadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""Mamba hybrid + UnifiedRadixCache."""
kl_threshold = 0.003
prefill_cache_assert = staticmethod(
make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE)
)
decode_cache_assert = staticmethod(
make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL)
)
@classmethod
def setUpClass(cls):
cls.model = MAMBA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--chunked-prefill-size",
"2048",
"--mem-fraction-static",
"0.85",
"--mamba-scheduler-strategy",
"extra_buffer",
"--mamba-track-interval",
str(MAMBA_TRACK_INTERVAL),
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""SWA hybrid + UnifiedRadixCache."""
kl_threshold = 0.03
gsm8k_threshold = 0.75
mmlu_threshold = 0.75
@classmethod
def setUpClass(cls):
cls.model = SWA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--mem-fraction-static",
"0.7",
"--disable-piecewise-cuda-graph",
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "0"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,775 @@
"""Large-scale benchmark + fuzz correctness tests for UnifiedRadixCache.
Usage (standalone):
bench: python3 test/registered/unit/mem_cache/test_unified_radix_cache_bench.py --num-seqs 5000 --verify --components mamba legacy-mamba swa legacy-swa
CI Test: python -m pytest test/registered/unit/mem_cache/test_unified_radix_cache_bench.py -v -s
"""
import argparse
import gc
import logging
import random
import statistics
import time
import unittest
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Callable
import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
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
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.mem_cache.unified_cache_components.tree_component import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, suite="stage-b-test-1-gpu-small")
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
_PAGE_SIZE = 1
_HEAD_NUM = 2
_HEAD_DIM = 16
_NUM_LAYERS = 8
_GLOBAL_INTERVAL = 4
_DTYPE = torch.bfloat16
_SWA_WINDOW_SIZE = 128
_BENCH_NUM_SEQS = 5000
_BENCH_KV_SIZE = 500_000
_BENCH_CHUNK_LEN = 256
_DEFAULT_COMPONENTS = (ComponentType.FULL, ComponentType.MAMBA)
@contextmanager
def _suppress_logs():
root = logging.getLogger()
prev = root.level
root.setLevel(logging.WARNING)
try:
yield
finally:
root.setLevel(prev)
def _full_attention_layer_ids():
return list(range(_GLOBAL_INTERVAL - 1, _NUM_LAYERS, _GLOBAL_INTERVAL))
def _non_full_layer_ids():
full = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full]
# ===================================================================
# Sequence generator
# ===================================================================
def gen_random_sequences(
num_seqs: int = 2000,
chunk_len: int = 256,
vocab_size: int = 32000,
seed: int = 42,
) -> list[list[int]]:
"""Generate *num_seqs* token sequences with tree-like prefix sharing.
Phase 1 (50%): chain growth — each new seq extends a random existing one.
Phase 2 (50%): fan-out burst — multiple children from the same parent.
"""
rng = random.Random(seed)
root_prefix = [rng.randint(1, vocab_size) for _ in range(max(1, chunk_len // 4))]
sequences: list[list[int]] = [root_prefix[:]]
# Phase 1: chain growth
for _ in range(num_seqs // 2):
parent = rng.choice(sequences)
sequences.append(
parent + [rng.randint(1, vocab_size)] * rng.randint(1, chunk_len)
)
# Phase 2: fan-out burst
remaining = num_seqs - num_seqs // 2
while remaining > 0:
fan = min(rng.randint(2, 10), remaining)
parent = rng.choice(sequences)
for _ in range(fan):
sequences.append(
parent + [rng.randint(1, vocab_size)] * rng.randint(1, chunk_len)
)
remaining -= fan
rng.shuffle(sequences)
return sequences
# ===================================================================
# Cache factory
# ===================================================================
def create_bench_cache(
kv_size,
max_num_reqs,
max_context_len,
components,
page_size=_PAGE_SIZE,
tree_cls=None,
):
"""Create cache. Returns (tree, allocator, req_to_token_pool, make_req)."""
device = get_device()
has_mamba = ComponentType.MAMBA in components
has_swa = ComponentType.SWA in components
mamba2_cache_params = None
if has_mamba:
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=256,
n_groups=1,
num_heads=2,
head_dim=16,
state_size=16,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_non_full_layer_ids()
)
# --- req_to_token pool ---
if has_mamba:
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=max(max_num_reqs * 2, 200),
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_non_full_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
else:
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
# --- KV pool + allocator ---
if has_swa:
from sglang.srt.mem_cache.swa_memory_pool import (
SWAKVPool,
SWATokenToKVPoolAllocator,
)
pool = SWAKVPool(
size=kv_size,
size_swa=kv_size,
page_size=page_size,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
swa_attention_layer_ids=_non_full_layer_ids(),
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
)
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size,
page_size=page_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
else:
pool = HybridLinearKVPool(
size=kv_size,
dtype=_DTYPE,
page_size=page_size,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
enable_memory_saver=False,
mamba_pool=req_to_token_pool.mamba_pool if has_mamba else None,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
# --- tree ---
if tree_cls is None:
tree_cls = UnifiedRadixCache
tree = tree_cls(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
tree_components=components if tree_cls is UnifiedRadixCache else None,
sliding_window_size=_SWA_WINDOW_SIZE if has_swa else None,
)
)
_rid = [0]
def make_req():
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.sampling.sampling_params import SamplingParams
req = Req(
rid=_rid[0],
origin_input_text="",
origin_input_ids=[],
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
)
_rid[0] += 1
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ===================================================================
# Shared bench environment + helpers
# ===================================================================
@dataclass
class _Env:
tree: object
alloc: object
rtp: object
make_req: Callable
seqs: list
has_mamba: bool
avg_tokens: int
def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None):
"""Create sequences + cache, return shared _Env."""
if components is None:
components = _DEFAULT_COMPONENTS
seqs = gen_random_sequences(num_seqs=num_seqs, chunk_len=chunk_len)
max_seq_len = max(len(s) for s in seqs)
avg_tokens = sum(len(s) for s in seqs) // len(seqs)
with _suppress_logs():
tree, alloc, rtp, make_req = create_bench_cache(
kv_size=kv_size,
max_num_reqs=num_seqs + 100,
max_context_len=max_seq_len + 10,
components=components,
tree_cls=tree_cls,
)
return _Env(
tree, alloc, rtp, make_req, seqs, ComponentType.MAMBA in components, avg_tokens
)
def _alloc_with_evict(env, n):
"""Alloc *n* tokens, evicting if necessary. Returns tensor or None."""
v = env.alloc.alloc(n)
if v is None:
env.tree.evict(EvictParams(num_tokens=n * 2, mamba_num=2))
v = env.alloc.alloc(n)
return v
def _insert_seq(env, seq):
"""Insert one sequence (alloc + evict-fallback). Returns True on success."""
v = _alloc_with_evict(env, len(seq))
if v is None:
return False
mamba_val = None
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))
return True
def _populate(env, count):
"""Insert first *count* sequences (with evict-fallback)."""
for seq in env.seqs[:count]:
_insert_seq(env, seq)
def _fill_no_evict(env):
"""Insert sequences until pool exhausted (no eviction). Returns count."""
inserted = 0
for seq in env.seqs:
v = env.alloc.alloc(len(seq))
if v is None:
break
mamba_val = None
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))
inserted += 1
return inserted
# ===================================================================
# Benchmark result + runner
# ===================================================================
@dataclass
class BenchResult:
name: str
num_ops: int
total_tokens: int
elapsed_s: float
latencies_us: list[float]
@property
def ops_per_sec(self):
return self.num_ops / self.elapsed_s if self.elapsed_s > 0 else 0
@property
def tokens_per_sec(self):
return self.total_tokens / self.elapsed_s if self.elapsed_s > 0 else 0
@property
def p50_us(self):
return statistics.median(self.latencies_us) if self.latencies_us else 0
@property
def p99_us(self):
if not self.latencies_us:
return 0
idx = int(len(self.latencies_us) * 0.99)
return sorted(self.latencies_us)[min(idx, len(self.latencies_us) - 1)]
def report(self):
tok = (
f"{self.tokens_per_sec:>12,.0f} tok/s"
if self.total_tokens > 0
else f"{'N/A':>12s} tok/s"
)
return (
f" {self.name:<18s} | {tok} | {self.ops_per_sec:>10,.0f} ops/s | "
f"p50={self.p50_us:>8,.0f}us p99={self.p99_us:>8,.0f}us"
)
def bench_api(
name, setup_fn, op_fn, num_ops, tokens_per_op=0, warmup=10, verify_fn=None
):
"""Time *op_fn(item)* for each item from *setup_fn()*.
*verify_fn*, if provided, runs during warmup and once after timing
(excluded from latency measurement).
"""
items = setup_fn()
assert (
len(items) >= num_ops + warmup
), f"need {num_ops + warmup} items, got {len(items)}"
for i in range(warmup):
op_fn(items[i])
if verify_fn:
verify_fn(items[i])
gc.collect()
gc_was = gc.isenabled()
gc.disable()
latencies: list[float] = []
t0 = time.perf_counter()
for i in range(warmup, warmup + num_ops):
ts = time.perf_counter()
op_fn(items[i])
latencies.append((time.perf_counter() - ts) * 1e6)
elapsed = time.perf_counter() - t0
if gc_was:
gc.enable()
if verify_fn:
verify_fn(items[warmup + num_ops - 1])
return BenchResult(
name,
num_ops,
tokens_per_op * num_ops if tokens_per_op > 0 else 0,
elapsed,
latencies,
)
# ===================================================================
# Five benchmark scenarios
# ===================================================================
def bench_insert(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Insert throughput (alloc + evict-fallback + insert)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
warmup = min(20, num_seqs // 10)
return bench_api(
"insert",
lambda: list(range(len(env.seqs))),
lambda idx: _insert_seq(env, env.seqs[idx]),
num_seqs - warmup,
env.avg_tokens,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_match_prefix(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Prefix matching throughput (hit / partial / miss mix)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
_populate(env, num_seqs // 2)
rng = random.Random(123)
pop = num_seqs // 2
queries: list[list[int]] = []
for _ in env.seqs:
roll = rng.random()
if roll < 0.33:
queries.append(env.seqs[rng.randint(0, pop - 1)])
elif roll < 0.66:
base = env.seqs[rng.randint(0, pop - 1)]
queries.append(base + [rng.randint(1, 32000)] * rng.randint(10, 100))
else:
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)))
assert len(r1.device_indices) == len(r2.device_indices), "match not idempotent"
warmup = min(20, len(queries) // 10)
return bench_api(
"match_prefix",
lambda: queries,
lambda q: env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q))),
min(len(queries) - warmup, num_seqs),
env.avg_tokens,
warmup,
verify_fn if verify else None,
)
def bench_evict(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Eviction throughput — fill pool then repeatedly evict batches."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
inserted = _fill_no_evict(env)
evict_batch = max(100, kv_size // 200)
num_evictions = max(inserted // 5, 100)
items = [(evict_batch,)] * (num_evictions + 50)
warmup = min(20, num_evictions // 10)
return bench_api(
"evict",
lambda: items,
lambda item: env.tree.evict(EvictParams(num_tokens=item[0], mamba_num=2)),
num_evictions - warmup,
evict_batch,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_lock_unlock(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Lock/unlock throughput — match nodes then cycle lock/unlock."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
_populate(env, num_seqs // 2)
nodes = []
for seq in env.seqs[: num_seqs // 2]:
r = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
if r.last_device_node != env.tree.root_node:
nodes.append(r.last_device_node)
if not nodes:
return BenchResult("lock_unlock", 0, 0, 0, [])
rng = random.Random(99)
num_pairs = min(len(nodes) * 2, num_seqs)
items = [rng.choice(nodes) for _ in range(num_pairs + 50)]
def op_fn(node):
lr = env.tree.inc_lock_ref(node)
env.tree.dec_lock_ref(
node,
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
)
warmup = min(20, num_pairs // 10)
return bench_api(
"lock_unlock",
lambda: items,
op_fn,
num_pairs - warmup,
0,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_cache_finished(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""cache_finished_req throughput — full request lifecycle.
Simulates: match_prefix → inc_lock_ref → alloc → fill req_to_token → cache_finished_req.
"""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
# Pre-build Req objects with token IDs filled into req_to_token
req_items: list = []
for seq in env.seqs:
key = RadixKey(seq)
mr = env.tree.match_prefix(MatchPrefixParams(key=key))
matched_len = len(mr.device_indices)
node = mr.last_device_node
lr = env.tree.inc_lock_ref(node)
remaining = len(seq) - matched_len
if remaining > 0:
v = _alloc_with_evict(env, remaining)
if v is None:
env.tree.dec_lock_ref(
node,
DecLockRefParams(
swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)
),
)
continue
kv_indices = torch.cat([mr.device_indices, v])
else:
kv_indices = mr.device_indices
req = env.make_req()
req.origin_input_ids = list(seq)
req.output_ids = []
req.fill_ids = list(seq)
req.last_node = node
req.cache_protected_len = matched_len
req.kv_committed_len = len(seq)
req.kv_committed_freed = False
if hasattr(lr, "swa_uuid_for_lock"):
req.swa_uuid_for_lock = lr.swa_uuid_for_lock
env.rtp.req_to_token[req.req_pool_idx, : len(kv_indices)] = kv_indices
req_items.append(req)
if not req_items:
return BenchResult("cache_finished", 0, 0, 0, [])
warmup = min(20, len(req_items) // 10)
return bench_api(
"cache_finished",
lambda: req_items,
lambda req: env.tree.cache_finished_req(req, is_insert=True),
len(req_items) - warmup,
env.avg_tokens,
warmup,
# Pool math doesn't hold here (many reqs still hold allocated tokens).
(lambda _: env.tree.sanity_check()) if verify else None,
)
# ===================================================================
# Runner
# ===================================================================
ALL_BENCHMARKS = {
"insert": bench_insert,
"match": bench_match_prefix,
"evict": bench_evict,
"lock": bench_lock_unlock,
"cache_finished": bench_cache_finished,
}
def run_all_benchmarks(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
benchmarks=None,
tree_cls=None,
):
if components is None:
components = _DEFAULT_COMPONENTS
if benchmarks is None or "all" in benchmarks:
benchmarks = list(ALL_BENCHMARKS.keys())
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
impl_name = (tree_cls or UnifiedRadixCache).__name__
results = []
for name in benchmarks:
if name not in ALL_BENCHMARKS:
print(f"[WARN] Unknown benchmark: {name}, skipping")
continue
results.append(
ALL_BENCHMARKS[name](
num_seqs=num_seqs,
chunk_len=chunk_len,
kv_size=kv_size,
components=components,
verify=verify,
tree_cls=tree_cls,
)
)
print("=" * 100)
print(
f"{impl_name} Benchmark | "
f"num_seqs={num_seqs} chunk_len={chunk_len} kv_size={kv_size} "
f"components={[c.value for c in components]} verify={verify}"
)
print("-" * 100)
for r in results:
print(r.report())
print("=" * 100)
return results
# ===================================================================
# pytest wrapper
# ===================================================================
class TestUnifiedRadixCacheBench(unittest.TestCase):
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def test_bench_insert(self):
r = bench_insert(_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
def test_bench_match_prefix(self):
r = bench_match_prefix(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
def test_bench_evict(self):
r = bench_evict(_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True)
self.assertGreater(r.num_ops, 0)
def test_bench_lock_unlock(self):
r = bench_lock_unlock(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
def test_bench_cache_finished(self):
r = bench_cache_finished(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
# ===================================================================
# CLI
# ===================================================================
_TREE_CONFIGS = {
"full": ((ComponentType.FULL,), None),
"mamba": ((ComponentType.FULL, ComponentType.MAMBA), None),
"swa": ((ComponentType.FULL, ComponentType.SWA), None),
"all": ((ComponentType.FULL, ComponentType.SWA, ComponentType.MAMBA), None),
"legacy-mamba": ((ComponentType.FULL, ComponentType.MAMBA), MambaRadixCache),
"legacy-swa": ((ComponentType.FULL, ComponentType.SWA), SWARadixCache),
}
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="UnifiedRadixCache benchmark")
parser.add_argument("--num-seqs", type=int, default=5000)
parser.add_argument("--chunk-len", type=int, default=256)
parser.add_argument("--kv-size", type=int, default=500_000)
parser.add_argument(
"--components",
nargs="+",
choices=list(_TREE_CONFIGS.keys()),
default=["mamba", "legacy-mamba"],
help="Component configs to benchmark",
)
parser.add_argument(
"--verify", action="store_true", help="Enable correctness assertions"
)
parser.add_argument(
"--benchmarks",
nargs="+",
default=["all"],
help="insert match evict lock cache_finished all",
)
args, _ = parser.parse_known_args()
for comp_name in args.components:
components, tree_cls = _TREE_CONFIGS[comp_name]
run_all_benchmarks(
num_seqs=args.num_seqs,
chunk_len=args.chunk_len,
kv_size=args.kv_size,
components=components,
verify=args.verify,
benchmarks=args.benchmarks,
tree_cls=tree_cls,
)
@@ -0,0 +1,883 @@
"""Unit tests for UnifiedRadixCache"""
import unittest
import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
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 (
DecLockRefParams,
EvictParams,
EvictResult,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import available_and_evictable_str
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.unified_cache_components.tree_component import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
)
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, suite="stage-b-test-1-gpu-small")
# ---------------------------------------------------------------------------
# Shared constants
# ---------------------------------------------------------------------------
_PAGE_SIZE = 1
_HEAD_NUM = 2
_HEAD_DIM = 128
_NUM_LAYERS = 24
_GLOBAL_INTERVAL = 4
_DTYPE = torch.bfloat16
def _full_attention_layer_ids():
return [i for i in range(_GLOBAL_INTERVAL - 1, _NUM_LAYERS, _GLOBAL_INTERVAL)]
def _mamba_layer_ids():
full_set = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full_set]
def _swa_attention_layer_ids():
full_set = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full_set]
# ===================================================================
# Test: Full + Mamba components (no SWA)
# ===================================================================
class TestUnifiedRadixCacheMamba(unittest.TestCase):
"""UnifiedRadixCache with (Full, Mamba) components."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
max_num_reqs: int = 10,
mamba_cache_size: int = 20,
max_context_len: int = 128,
):
device = get_device()
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=4096,
n_groups=16,
num_heads=32,
head_dim=128,
state_size=128,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_mamba_layer_ids()
)
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_mamba_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
pool = HybridLinearKVPool(
size=kv_size,
dtype=_DTYPE,
page_size=_PAGE_SIZE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
enable_memory_saver=False,
mamba_pool=req_to_token_pool.mamba_pool,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
tree_components=(ComponentType.FULL, ComponentType.MAMBA),
),
)
def make_req():
sp = SamplingParams(temperature=0, max_new_tokens=1)
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
sampling_params=sp,
)
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ------- insert + match -------
def test_insert_and_match_basic(self):
tree, alloc, _, make_req = self._build_tree()
# Insert [1,2,3]
req1 = make_req()
v1 = alloc.alloc(3)
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=v1,
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(result.prefix_len, 0)
# Insert [1,2,3,4,5] — shares prefix [1,2,3]
req2 = make_req()
v2 = alloc.alloc(5)
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=v2,
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(result.prefix_len, 3)
# Match [1,2,3,4,5] — full hit
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 5)
# Match [1,2,3,4,5,6] — partial hit (5 tokens)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5, 6])))
self.assertEqual(len(m.device_indices), 5)
# Match [10,11] — no hit
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([10, 11])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- evict: full-only -------
def test_evict_full_tokens(self):
tree, alloc, _, make_req = self._build_tree()
# Insert two disjoint sequences
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12]),
value=alloc.alloc(3),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 6)
# Evict 3 full tokens — should remove one leaf
result = tree.evict(EvictParams(num_tokens=3))
self.assertIsInstance(result, EvictResult)
self.assertGreaterEqual(result.num_tokens_evicted, 3)
self.assertTrue(tree.full_evictable_size() <= 3)
tree.sanity_check()
# ------- evict: mamba-only -------
def test_evict_mamba_only(self):
tree, alloc, rtp, make_req = self._build_tree()
mamba_pool = rtp.mamba_pool
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7]),
value=alloc.alloc(7),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.mamba_evictable_size(), 2)
# Evict 1 mamba state
result = tree.evict(EvictParams(num_tokens=0, mamba_num=1))
self.assertGreaterEqual(result.mamba_num_evicted, 1)
# After mamba eviction on an internal node, full tokens remain
self.assertGreaterEqual(tree.full_evictable_size(), 0)
tree.sanity_check()
# ------- evict: mamba → match stops at tombstone -------
def test_evict_mamba_breaks_match(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Evict all mamba (2 states)
tree.evict(EvictParams(num_tokens=0, mamba_num=2))
self.assertEqual(tree.mamba_evictable_size(), 0)
# Now match should return 0 because mamba validator fails
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- evict: lock_ref protection -------
def test_evict_respects_lock_ref(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12]),
value=alloc.alloc(3),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Lock the first leaf
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
locked_node = m.last_device_node
tree.inc_lock_ref(locked_node)
# Evict all full tokens — only unlocked leaf should be evicted
result = tree.evict(EvictParams(num_tokens=6))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
# [1,2,3] is still matchable because it was locked
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
self.assertEqual(len(m.device_indices), 3)
# Unlock and verify we can now evict it
tree.dec_lock_ref(locked_node)
result = tree.evict(EvictParams(num_tokens=3))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
tree.sanity_check()
# ------- evict: verify EvictResult accounting -------
def test_evict_result_accounting(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Request 0 mamba + 3 full → full evicted, mamba cascaded
result = tree.evict(EvictParams(num_tokens=3))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
# Leaf eviction cascades all components; mamba also freed
self.assertGreaterEqual(result.mamba_num_evicted, 1)
tree.sanity_check()
# ------- insert: prev_prefix_len controls overlap free -------
def test_insert_prev_prefix_len(self):
tree, alloc, _, make_req = self._build_tree()
initial_avail = alloc.available_size()
# Step 1: Insert [1,2,3]
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(alloc.available_size(), initial_avail - 3)
# Step 2: Insert [1,2,3,4,5] with prev_prefix_len=0 → frees overlap [0:3]
req2 = make_req()
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=0,
)
)
self.assertEqual(result.prefix_len, 3)
# alloc 5, freed 3 overlap, stored 2 new → net -2
self.assertEqual(alloc.available_size(), initial_avail - 3 - 2)
# Step 3: Insert [1,2,3,4,5,6] with prev_prefix_len=5 → nothing freed
req3 = make_req()
avail_before = alloc.available_size()
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6]),
value=alloc.alloc(6),
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=5,
)
)
self.assertEqual(result.prefix_len, 5)
# alloc 6, freed 0, stored 1 → net -6
self.assertEqual(alloc.available_size(), avail_before - 6)
tree.sanity_check()
# ------- available_and_evictable_str + pretty_print -------
def test_diagnostics(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
diag = tree.available_and_evictable_str()
self.assertIn("Available full tokens", diag)
self.assertIn("mamba", diag.lower())
diag2 = available_and_evictable_str(tree)
self.assertIn("Available full tokens", diag2)
tree.pretty_print()
tree.sanity_check()
# ===================================================================
# Test: Full + SWA + Mamba components
# ===================================================================
class TestUnifiedRadixCacheSWAMamba(unittest.TestCase):
"""UnifiedRadixCache with (Full, SWA, Mamba) components — the most complex config."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
kv_size_swa: int = 64,
max_num_reqs: int = 10,
mamba_cache_size: int = 20,
max_context_len: int = 128,
sliding_window_size: int = 4,
):
device = get_device()
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=4096,
n_groups=16,
num_heads=32,
head_dim=128,
state_size=128,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_mamba_layer_ids()
)
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_mamba_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
kv_pool = SWAKVPool(
size=kv_size,
size_swa=kv_size_swa,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
swa_attention_layer_ids=_swa_attention_layer_ids(),
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
)
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size_swa,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
device=device,
kvcache=kv_pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
sliding_window_size=sliding_window_size,
tree_components=(
ComponentType.FULL,
ComponentType.SWA,
ComponentType.MAMBA,
),
),
)
def make_req():
sp = SamplingParams(temperature=0, max_new_tokens=1)
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
sampling_params=sp,
)
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ------- basic insert + match with SWA -------
def test_insert_and_match_with_swa(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Match: SWA validator requires contiguous window >= sliding_window_size
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
# With sliding_window_size=4 and 5 tokens on single node → should match
self.assertEqual(len(m.device_indices), 5)
tree.sanity_check()
# ------- evict SWA → cascade Mamba -------
def test_evict_swa_cascades_mamba(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
# Build tree: [1,2,3] → [4,5,6,7]
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7]),
value=alloc.alloc(7),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
initial_mamba = tree.mamba_evictable_size()
# Evict SWA — on internal node, SWA eviction cascades to Mamba (priority: swa=1 > mamba=0)
result = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=3))
self.assertGreaterEqual(result.swa_num_tokens_evicted, 0)
tree.sanity_check()
# ------- evict full leaf -------
def test_evict_full_leaf_cascades_all(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12, 13, 14]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 10)
# Evict one leaf (5 full tokens) → also cascades SWA + Mamba
result = tree.evict(EvictParams(num_tokens=5))
self.assertGreaterEqual(result.num_tokens_evicted, 5)
# Leaf eviction should cascade all components
self.assertGreaterEqual(result.mamba_num_evicted, 1)
self.assertGreaterEqual(result.swa_num_tokens_evicted, 0)
tree.sanity_check()
# ------- evict with SWA lock -------
def test_swa_lock_protects_from_eviction(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12, 13, 14]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Lock the first entry
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
lock_result = tree.inc_lock_ref(m.last_device_node)
# Try to evict all full tokens
result = tree.evict(EvictParams(num_tokens=10))
# Only the unlocked one (5 tokens) should be evictable
self.assertGreaterEqual(result.num_tokens_evicted, 5)
# Locked one is still matchable
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 5)
# Unlock
tree.dec_lock_ref(
m.last_device_node,
DecLockRefParams(swa_uuid_for_lock=lock_result.swa_uuid_for_lock),
)
tree.sanity_check()
# ------- cache_finished_req (with insert) -------
def test_cache_finished_req_insert(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3, 4, 5]
req.output_ids = [6, 7]
kv_len = len(req.origin_input_ids) + len(req.output_ids)
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.mamba_last_track_seqlen = kv_len
req.fill_ids = req.origin_input_ids + req.output_ids
tree.cache_finished_req(req, is_insert=True)
# Verify the tokens are in the tree
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5, 6, 7])))
self.assertEqual(len(m.device_indices), 7)
tree.sanity_check()
# ------- cache_finished_req (no insert) -------
def test_cache_finished_req_no_insert(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3]
req.output_ids = []
kv_len = 3
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.fill_ids = req.origin_input_ids
avail_before = alloc.available_size()
tree.cache_finished_req(req, is_insert=False)
# KV indices should be freed back
self.assertEqual(alloc.available_size(), avail_before + kv_len)
# Nothing in tree
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- cache_unfinished_req -------
def test_cache_unfinished_req(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3, 4, 5]
req.output_ids = []
req.fill_ids = req.origin_input_ids[:]
kv_len = len(req.fill_ids)
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.mamba_last_track_seqlen = kv_len
tree.cache_unfinished_req(req)
# After caching, prefix_indices should be set
self.assertGreater(len(req.prefix_indices), 0)
self.assertEqual(req.cache_protected_len, len(req.prefix_indices))
self.assertIsNotNone(req.last_node)
# Release the lock acquired by cache_unfinished_req before idle check
tree.dec_lock_ref(
req.last_node,
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
)
tree.sanity_check()
# ------- evict empty tree → no crash -------
def test_evict_empty_tree(self):
tree, alloc, _, _ = self._build_tree()
result = tree.evict(EvictParams(num_tokens=10, mamba_num=5))
self.assertEqual(result.num_tokens_evicted, 0)
self.assertEqual(result.mamba_num_evicted, 0)
tree.sanity_check()
# ------- multiple evictions until empty -------
def test_evict_until_empty(self):
tree, alloc, _, make_req = self._build_tree()
for i in range(5):
req = make_req()
tokens = list(range(i * 10, i * 10 + 5))
tree.insert(
InsertParams(
key=RadixKey(tokens),
value=alloc.alloc(5),
mamba_value=req.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 25)
# Evict all
result = tree.evict(EvictParams(num_tokens=100))
self.assertGreaterEqual(result.num_tokens_evicted, 25)
self.assertEqual(tree.full_evictable_size(), 0)
self.assertEqual(tree.mamba_evictable_size(), 0)
# Verify tree is empty (no matches)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([0, 1, 2, 3, 4])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- cow mamba on match -------
def test_match_cow_mamba(self):
tree, alloc, rtp, make_req = self._build_tree()
mamba_pool = rtp.mamba_pool
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Match with cow_mamba
req2 = make_req()
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5]), cow_mamba=True, req=req2)
)
self.assertEqual(len(m.device_indices), 5)
# req2 should now have its own mamba state (copied)
self.assertIsNotNone(req2.mamba_pool_idx)
# Verify the copy matches
src_value = m.last_device_node.component_data[ComponentType.MAMBA].value
self.assertTrue(
torch.all(
mamba_pool.mamba_cache.conv[0][:, req2.mamba_pool_idx]
== mamba_pool.mamba_cache.conv[0][:, src_value]
)
)
tree.sanity_check()
# ===================================================================
# Test: Helper functions
# ===================================================================
class TestUnifiedRadixCacheHelpers(unittest.TestCase):
"""Tests for internal helper functions of UnifiedRadixCache."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
max_num_reqs: int = 10,
max_context_len: int = 128,
):
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
device = get_device()
req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
kv_pool = MHATokenToKVPool(
size=kv_size,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
layer_num=_NUM_LAYERS,
device=device,
enable_memory_saver=False,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=kv_pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
tree_components=(
ComponentType.FULL,
), # Full attention only, no mamba/swa
),
)
return tree, allocator
def test_readonly_does_not_modify_tree(self):
"""Verify readonly match does not modify tree structure (no split)."""
tree, alloc = self._build_tree()
# Insert [1, 2, 3, 4, 5]
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
)
)
def count_nodes(node):
count = 1
for child in node.children.values():
count += count_nodes(child)
return count
node_count_before = count_nodes(tree.root_node)
self.assertEqual(node_count_before, 2) # root_node and [1, 2, 3, 4, 5]
# Regular match with partial key [1, 2] creates a split
value, best_node, best_value_len = tree._match_prefix_helper(RadixKey([1, 2]))
# Regular match with partial key [1, 2, 3, 4] creates a split
value, best_node, best_value_len = tree._match_prefix_helper(
RadixKey([1, 2, 3, 4])
)
self.assertEqual(best_value_len, 2)
self.assertEqual(best_node.key.token_ids, [3, 4])
node_count_after_regular = count_nodes(tree.root_node)
self.assertEqual(node_count_after_regular, node_count_before + 2)
# Readonly match with partial key [1, 2, 3] should NOT create a split
value, best_node, best_value_len = tree._match_prefix_helper_readonly(
RadixKey([1, 2, 3])
)
self.assertEqual(best_value_len, 1)
self.assertEqual(best_node.key.token_ids, [1, 2])
node_count_after_readonly = count_nodes(tree.root_node)
self.assertEqual(node_count_after_readonly, node_count_after_regular)
tree.sanity_check()
if __name__ == "__main__":
unittest.main()