Add page_size and SWA coverage to unified radix cache bench test (#22815)
This commit is contained in:
@@ -38,12 +38,11 @@ from sglang.srt.server_args import ServerArgs, set_global_server_args_for_schedu
|
||||
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")
|
||||
register_cuda_ci(est_time=120, suite="stage-b-test-1-gpu-small")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
# ---------------------------------------------------------------------------
|
||||
_PAGE_SIZE = 1
|
||||
_HEAD_NUM = 2
|
||||
_HEAD_DIM = 16
|
||||
_NUM_LAYERS = 8
|
||||
@@ -126,8 +125,9 @@ def create_bench_cache(
|
||||
max_num_reqs,
|
||||
max_context_len,
|
||||
components,
|
||||
page_size=_PAGE_SIZE,
|
||||
page_size=1,
|
||||
tree_cls=None,
|
||||
sliding_window_size=_SWA_WINDOW_SIZE,
|
||||
):
|
||||
"""Create cache. Returns (tree, allocator, req_to_token_pool, make_req)."""
|
||||
device = get_device()
|
||||
@@ -161,7 +161,7 @@ def create_bench_cache(
|
||||
enable_memory_saver=False,
|
||||
cache_params=mamba2_cache_params,
|
||||
mamba_layer_ids=_non_full_layer_ids(),
|
||||
enable_mamba_extra_buffer=False,
|
||||
enable_mamba_extra_buffer=(page_size > 1),
|
||||
speculative_num_draft_tokens=3,
|
||||
)
|
||||
else:
|
||||
@@ -233,7 +233,7 @@ def create_bench_cache(
|
||||
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,
|
||||
sliding_window_size=sliding_window_size if has_swa else None,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -267,10 +267,12 @@ class _Env:
|
||||
make_req: Callable
|
||||
seqs: list
|
||||
has_mamba: bool
|
||||
has_swa: bool
|
||||
page_size: int
|
||||
avg_tokens: int
|
||||
|
||||
|
||||
def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None):
|
||||
def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None, page_size=1):
|
||||
"""Create sequences + cache, return shared _Env."""
|
||||
if components is None:
|
||||
components = _DEFAULT_COMPONENTS
|
||||
@@ -283,19 +285,44 @@ def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None):
|
||||
max_num_reqs=num_seqs + 100,
|
||||
max_context_len=max_seq_len + 10,
|
||||
components=components,
|
||||
page_size=page_size,
|
||||
tree_cls=tree_cls,
|
||||
)
|
||||
return _Env(
|
||||
tree, alloc, rtp, make_req, seqs, ComponentType.MAMBA in components, avg_tokens
|
||||
tree,
|
||||
alloc,
|
||||
rtp,
|
||||
make_req,
|
||||
seqs,
|
||||
ComponentType.MAMBA in components,
|
||||
ComponentType.SWA in components,
|
||||
page_size,
|
||||
avg_tokens,
|
||||
)
|
||||
|
||||
|
||||
def _alloc(env, n):
|
||||
if env.has_swa and env.page_size > 1:
|
||||
ps = env.page_size
|
||||
aligned = ((n + ps - 1) // ps) * ps
|
||||
if aligned > env.alloc.full_attn_allocator.available_size():
|
||||
return None
|
||||
if aligned > env.alloc.swa_attn_allocator.available_size():
|
||||
return None
|
||||
full_indices = env.alloc.full_attn_allocator.alloc(aligned)
|
||||
swa_indices = env.alloc.swa_attn_allocator.alloc(aligned)
|
||||
assert full_indices is not None and swa_indices is not None
|
||||
env.alloc.full_to_swa_index_mapping[full_indices] = swa_indices
|
||||
return full_indices[:n]
|
||||
return env.alloc.alloc(n)
|
||||
|
||||
|
||||
def _alloc_with_evict(env, n):
|
||||
"""Alloc *n* tokens, evicting if necessary. Returns tensor or None."""
|
||||
v = env.alloc.alloc(n)
|
||||
v = _alloc(env, n)
|
||||
if v is None:
|
||||
env.tree.evict(EvictParams(num_tokens=n * 2, mamba_num=2))
|
||||
v = env.alloc.alloc(n)
|
||||
v = _alloc(env, n)
|
||||
return v
|
||||
|
||||
|
||||
@@ -322,7 +349,7 @@ 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))
|
||||
v = _alloc(env, len(seq))
|
||||
if v is None:
|
||||
break
|
||||
mamba_val = None
|
||||
@@ -430,9 +457,10 @@ def bench_insert(
|
||||
components=None,
|
||||
verify=False,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
"""Insert throughput (alloc + evict-fallback + insert)."""
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size)
|
||||
warmup = min(20, num_seqs // 10)
|
||||
|
||||
return bench_api(
|
||||
@@ -453,9 +481,10 @@ def bench_match_prefix(
|
||||
components=None,
|
||||
verify=False,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
"""Prefix matching throughput (hit / partial / miss mix)."""
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size)
|
||||
_populate(env, num_seqs // 2)
|
||||
|
||||
rng = random.Random(123)
|
||||
@@ -495,9 +524,10 @@ def bench_evict(
|
||||
components=None,
|
||||
verify=False,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
"""Eviction throughput — fill pool then repeatedly evict batches."""
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size)
|
||||
inserted = _fill_no_evict(env)
|
||||
|
||||
evict_batch = max(100, kv_size // 200)
|
||||
@@ -523,9 +553,10 @@ def bench_lock_unlock(
|
||||
components=None,
|
||||
verify=False,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
"""Lock/unlock throughput — match nodes then cycle lock/unlock."""
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size)
|
||||
_populate(env, num_seqs // 2)
|
||||
|
||||
nodes = []
|
||||
@@ -566,12 +597,13 @@ def bench_cache_finished(
|
||||
components=None,
|
||||
verify=False,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
"""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)
|
||||
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size)
|
||||
|
||||
# Pre-build Req objects with token IDs filled into req_to_token
|
||||
req_items: list = []
|
||||
@@ -646,6 +678,7 @@ def run_all_benchmarks(
|
||||
verify=False,
|
||||
benchmarks=None,
|
||||
tree_cls=None,
|
||||
page_size=1,
|
||||
):
|
||||
if components is None:
|
||||
components = _DEFAULT_COMPONENTS
|
||||
@@ -653,7 +686,7 @@ def run_all_benchmarks(
|
||||
benchmarks = list(ALL_BENCHMARKS.keys())
|
||||
|
||||
set_global_server_args_for_scheduler(
|
||||
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
|
||||
ServerArgs(model_path="dummy", page_size=page_size)
|
||||
)
|
||||
|
||||
impl_name = (tree_cls or UnifiedRadixCache).__name__
|
||||
@@ -670,6 +703,7 @@ def run_all_benchmarks(
|
||||
components=components,
|
||||
verify=verify,
|
||||
tree_cls=tree_cls,
|
||||
page_size=page_size,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -677,7 +711,7 @@ def run_all_benchmarks(
|
||||
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}"
|
||||
f"page_size={page_size} components={[c.value for c in components]} verify={verify}"
|
||||
)
|
||||
print("-" * 100)
|
||||
for r in results:
|
||||
@@ -689,42 +723,99 @@ def run_all_benchmarks(
|
||||
# ===================================================================
|
||||
# pytest wrapper
|
||||
# ===================================================================
|
||||
class TestUnifiedRadixCacheBench(unittest.TestCase):
|
||||
_CI_BENCH_CONFIGS = [
|
||||
dict(
|
||||
label="FULL_MAMBA_ps1",
|
||||
components=(ComponentType.FULL, ComponentType.MAMBA),
|
||||
page_size=1,
|
||||
num_seqs=5000,
|
||||
kv_size=500_000,
|
||||
),
|
||||
dict(
|
||||
label="FULL_SWA_ps1",
|
||||
components=(ComponentType.FULL, ComponentType.SWA),
|
||||
page_size=1,
|
||||
num_seqs=1000,
|
||||
kv_size=100_000,
|
||||
),
|
||||
dict(
|
||||
label="FULL_ps16",
|
||||
components=(ComponentType.FULL,),
|
||||
page_size=16,
|
||||
num_seqs=1000,
|
||||
kv_size=100_000,
|
||||
),
|
||||
dict(
|
||||
label="FULL_SWA_ps16",
|
||||
components=(ComponentType.FULL, ComponentType.SWA),
|
||||
page_size=16,
|
||||
num_seqs=1000,
|
||||
kv_size=100_000,
|
||||
),
|
||||
dict(
|
||||
label="FULL_ps128",
|
||||
components=(ComponentType.FULL,),
|
||||
page_size=128,
|
||||
num_seqs=1000,
|
||||
kv_size=200_000,
|
||||
),
|
||||
dict(
|
||||
label="FULL_SWA_ps128",
|
||||
components=(ComponentType.FULL, ComponentType.SWA),
|
||||
page_size=128,
|
||||
num_seqs=1000,
|
||||
kv_size=200_000,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class _BenchSuite:
|
||||
"""Mixin: subclass must set bench_cfg dict with keys: label, components, page_size, num_seqs, kv_size."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
set_global_server_args_for_scheduler(
|
||||
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
|
||||
ServerArgs(model_path="dummy", page_size=cls.bench_cfg["page_size"])
|
||||
)
|
||||
|
||||
def _run(self, bench_fn):
|
||||
cfg = self.bench_cfg
|
||||
r = bench_fn(
|
||||
cfg["num_seqs"],
|
||||
_BENCH_CHUNK_LEN,
|
||||
cfg["kv_size"],
|
||||
components=cfg["components"],
|
||||
verify=True,
|
||||
page_size=cfg["page_size"],
|
||||
)
|
||||
self.assertGreater(r.num_ops, 0)
|
||||
self.assertGreater(r.ops_per_sec, 0)
|
||||
|
||||
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)
|
||||
self._run(bench_insert)
|
||||
|
||||
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)
|
||||
self._run(bench_match_prefix)
|
||||
|
||||
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)
|
||||
self._run(bench_evict)
|
||||
|
||||
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)
|
||||
self._run(bench_lock_unlock)
|
||||
|
||||
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)
|
||||
self._run(bench_cache_finished)
|
||||
|
||||
|
||||
for _cfg in _CI_BENCH_CONFIGS:
|
||||
_name = f"TestBench_{_cfg['label']}"
|
||||
globals()[_name] = type(
|
||||
_name,
|
||||
(_BenchSuite, unittest.TestCase),
|
||||
{"bench_cfg": _cfg},
|
||||
)
|
||||
globals()[_name].__module__ = __name__
|
||||
del _cfg, _name
|
||||
|
||||
|
||||
# ===================================================================
|
||||
@@ -751,6 +842,7 @@ if __name__ == "__main__":
|
||||
default=["mamba", "legacy-mamba"],
|
||||
help="Component configs to benchmark",
|
||||
)
|
||||
parser.add_argument("--page-size", type=int, default=1)
|
||||
parser.add_argument(
|
||||
"--verify", action="store_true", help="Enable correctness assertions"
|
||||
)
|
||||
@@ -772,4 +864,5 @@ if __name__ == "__main__":
|
||||
verify=args.verify,
|
||||
benchmarks=args.benchmarks,
|
||||
tree_cls=tree_cls,
|
||||
page_size=args.page_size,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user