Remove swa and mamba radix cache (#40313)

This commit is contained in:
Ke Bao
2026-09-20 16:16:27 +08:00
committed by GitHub
parent 5c69e32abe
commit a8a4d86be9
24 changed files with 69 additions and 5184 deletions
+4 -6
View File
@@ -72,12 +72,10 @@ def handle_mamba_backend(server_args: Any):
def handle_int8_mamba_checkpoint(server_args: Any): def handle_int8_mamba_checkpoint(server_args: Any):
# The int8 mamba checkpoint pool is only wired into the built-in # The host-offload path (enabled by --enable-hierarchical-cache) and
# MambaRadixCache. The host-offload path (enabled by # custom radix-cache backends are NOT int8-aware: they would read int8
# --enable-hierarchical-cache) and custom radix-cache backends are NOT # checkpoint slots as bf16 active slots (wrong pool / out-of-range).
# int8-aware: they would read int8 checkpoint slots as bf16 active slots # Reject the combination up front rather than silently corrupting state.
# (wrong pool / out-of-range). Reject the combination up front rather than
# silently corrupting state.
cfg = resolving_view(server_args) cfg = resolving_view(server_args)
if not cfg.enable_int8_mamba_checkpoint: if not cfg.enable_int8_mamba_checkpoint:
return return
@@ -25,7 +25,7 @@ def _qwen4_exp_overrides(server_args: Any, hf_config: Any) -> dict:
"""Compressed QSA must own ``page_size`` here, """Compressed QSA must own ``page_size`` here,
so the qwen3_5 hybrid attention-shape policy is restated rather than shared. so the qwen3_5 hybrid attention-shape policy is restated rather than shared.
page_size=64 needs page-aligned full-KV allocation (slots are full_slot // ratio), page_size=64 needs page-aligned full-KV allocation (slots are full_slot // ratio),
which MambaRadixCache allows only with mamba extra-buffer or --disable-radix-cache. which in turn needs the mamba extra-buffer strategy or --disable-radix-cache.
""" """
cfg = resolving_view(server_args) cfg = resolving_view(server_args)
if ( if (
-3
View File
@@ -678,9 +678,6 @@ class Envs:
SGLANG_ENABLE_UNIFIED_RADIX_TREE = EnvBool(False) SGLANG_ENABLE_UNIFIED_RADIX_TREE = EnvBool(False)
# Registered TreeCore backend serving the unified radix cache. # Registered TreeCore backend serving the unified radix cache.
SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("python") SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("python")
# TODO(DSV4): @ispobock this has bug on main branch when retract
SGLANG_OPT_SWA_RADIX_CACHE_COMPACT = EnvBool(False)
SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT = EnvBool(False)
SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW = EnvBool(False) SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW = EnvBool(False)
# =================================================================== # ===================================================================
@@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Any
import torch import torch
from sglang.srt.mem_cache.radix_cache import RadixCache from sglang.srt.mem_cache.radix_cache import RadixCache
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import ( from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import (
RadixCacheWalkResult, RadixCacheWalkResult,
) )
@@ -34,7 +33,7 @@ def walk_radix_cache_for_canary(
return radix_cache.tree_core.walk_for_kv_canary( return radix_cache.tree_core.walk_for_kv_canary(
unlocked_only=unlocked_only, swa_resident_only=swa_resident_only unlocked_only=unlocked_only, swa_resident_only=swa_resident_only
) )
if cache_type is not RadixCache and cache_type is not SWARadixCache: if cache_type is not RadixCache:
raise NotImplementedError( raise NotImplementedError(
f"walk_radix_cache_for_canary does not support {cache_type.__name__}" f"walk_radix_cache_for_canary does not support {cache_type.__name__}"
) )
@@ -133,9 +132,6 @@ def _node_is_unlocked_for_canary(
if type(radix_cache) is RadixCache: if type(radix_cache) is RadixCache:
return node.lock_ref == 0 return node.lock_ref == 0
if type(radix_cache) is SWARadixCache:
return node.full_lock_ref == 0
raise NotImplementedError( raise NotImplementedError(
f"walk_radix_cache_for_canary does not support {type(radix_cache).__name__}" f"walk_radix_cache_for_canary does not support {type(radix_cache).__name__}"
) )
@@ -146,7 +142,5 @@ def _node_is_swa_resident_for_canary(
node: TreeNode, node: TreeNode,
radix_cache: BasePrefixCache, radix_cache: BasePrefixCache,
) -> bool: ) -> bool:
if type(radix_cache) is SWARadixCache: # RadixCache has no SWA tier, so every node it holds is resident.
return not node.swa_tombstone
return True return True
@@ -99,8 +99,8 @@ class SchedulerInvariantChecker:
session_held = self.pool_stats_observer.session_held_full_tokens() session_held = self.pool_stats_observer.session_held_full_tokens()
total = ps.full_capacity total = ps.full_capacity
elif self.is_hybrid_ssm: elif self.is_hybrid_ssm:
# Branch on cache type for the protected accessor (MambaRadixCache # Branch on cache type for the protected accessor (a mamba-capable
# splits full/mamba; ChunkCache only has the single protected_size). # cache splits full/mamba; ChunkCache only has the single protected_size).
# Use the allocator's `.size` for `total`: static max_total_num_tokens for # Use the allocator's `.size` for `total`: static max_total_num_tokens for
# non-unified pools, the dynamic byte-coordinated cap (matching # non-unified pools, the dynamic byte-coordinated cap (matching
# `available_size`) for the unified pool. # `available_size`) for the unified pool.
+1 -1
View File
@@ -38,7 +38,7 @@ to keep. The layout is specified in
Two groups sit outside that stack: Two groups sit outside that stack:
- **Radix cache** is its own axis. The per-model variants (`radix_cache.py`, - **Radix cache** is its own axis. The per-model variants (`radix_cache.py`,
`swa_radix_cache.py`, `mamba_radix_cache.py`, `hiradix_cache.py`, `chunk_cache.py`) `hiradix_cache.py`, `chunk_cache.py`)
are converging onto the **Unified Radix Cache** (`unified_cache/`, are converging onto the **Unified Radix Cache** (`unified_cache/`,
[#20415](https://github.com/sgl-project/sglang/issues/20415)), whose Full/SWA/Mamba [#20415](https://github.com/sgl-project/sglang/issues/20415)), whose Full/SWA/Mamba
component model is documented in component model is documented in
+2 -2
View File
@@ -271,10 +271,10 @@ def retraction_discard(req: Req, tree_cache: BasePrefixCache, backend: str) -> N
def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True): def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True):
assert (not req.kv.holds_kv) == req.kv.is_kv_released assert (not req.kv.holds_kv) == req.kv.is_kv_released
# MambaRadixCache may alloc mamba state before alloc KV cache # A mamba-capable cache may alloc mamba state before alloc KV cache
if not req.kv.holds_kv: if not req.kv.holds_kv:
assert tree_cache.supports_mamba(), ( assert tree_cache.supports_mamba(), (
"Only MambaRadixCache allow freeing before alloc" "Only a mamba-capable tree cache allows freeing before alloc"
) )
# TODO (csy, hanming): clean up this early allocation logic # TODO (csy, hanming): clean up this early allocation logic
if req.kv.holds_mamba: if req.kv.holds_mamba:
File diff suppressed because it is too large Load Diff
@@ -52,8 +52,8 @@ class PureSWARadixCache(RadixCache):
return 0 return 0
def sanity_check(self): def sanity_check(self):
"""No-op: PureSWARadixCache uses RadixCache's simple tree structure """No-op: an all-SWA model has no full tier, so there is no full/SWA
which doesn't need the dual-LRU sanity checks of SWARadixCache.""" split to cross-check."""
pass pass
def evict(self, params: EvictParams) -> EvictResult: def evict(self, params: EvictParams) -> EvictResult:
+2 -2
View File
@@ -74,8 +74,8 @@ class QSATokenToKVPool(HybridLinearKVPool):
"compressed QSA requires a paged full-KV cache with the page " "compressed QSA requires a paged full-KV cache with the page "
"a multiple of the compress ratio (compressed slots are " "a multiple of the compress ratio (compressed slots are "
f"full_slot // ratio): page_size={page_size}, " f"full_slot // ratio): page_size={page_size}, "
f"ratio={qsa_compress_ratio}. With MambaRadixCache this " f"ratio={qsa_compress_ratio}. This needs the mamba "
"needs the mamba extra-buffer strategy or " "extra-buffer strategy or "
"--disable-radix-cache (see the Qwen4-Exp arg overrides)." "--disable-radix-cache (see the Qwen4-Exp arg overrides)."
) )
# super().__init__ computes mem_usage via the overridden get_kv_size_bytes, # super().__init__ computes mem_usage via the overridden get_kv_size_bytes,
File diff suppressed because it is too large Load Diff
@@ -4,7 +4,7 @@ A component-based, pluggable prefix cache framework for SGLang that unifies Full
## Design Goals ## Design Goals
1. **Unified tree structure** — One radix tree manages all KV cache types instead of separate specialized implementations (`SWARadixCache`, `MambaRadixCache`, etc.). 1. **Unified tree structure** — One radix tree manages all KV cache types, replacing the separate specialized implementations that preceded it.
2. **Pluggable components** — Each attention/state type (Full, SWA, Mamba) is a `TreeComponent` that implements hook interfaces. Adding a new cache type only requires adding a new component. 2. **Pluggable components** — Each attention/state type (Full, SWA, Mamba) is a `TreeComponent` that implements hook interfaces. Adding a new cache type only requires adding a new component.
3. **Per-component resource isolation** — Each component has its own lock reference counting, evictable/protected size tracking, and eviction driver. Auxiliary components use per-component LRUs; Full uses device/host leaf sets. 3. **Per-component resource isolation** — Each component has its own lock reference counting, evictable/protected size tracking, and eviction driver. Auxiliary components use per-component LRUs; Full uses device/host leaf sets.
4. **Cascade eviction with priority** — When a component evicts a node, lower-or-equal-priority components on the same node are evicted together, maintaining cross-component consistency. 4. **Cascade eviction with priority** — When a component evicts a node, lower-or-equal-priority components on the same node are evicted together, maintaining cross-component consistency.
@@ -537,8 +537,7 @@ class MambaComponent(TreeComponent):
# slot's unflushed ring depth (`write_pos`), so on request finish cap # slot's unflushed ring depth (`write_pos`), so on request finish cap
# the donate to the last flush boundary (where temporal is current) # the donate to the last flush boundary (where temporal is current)
# and reset the cursor, keeping the donated checkpoint consistent with # and reset the cursor, keeping the donated checkpoint consistent with
# its key length. page_size is asserted == 1, so no realign. Mirrors # its key length. page_size is asserted == 1, so no realign.
# MambaRadixCache.cache_finished_req.
if is_finished: if is_finished:
write_pos_buf = ( write_pos_buf = (
self.cache.req_to_token_pool.mamba_pool.replayssm_write_pos self.cache.req_to_token_pool.mamba_pool.replayssm_write_pos
+1 -1
View File
@@ -4611,7 +4611,7 @@ def get_extend_input_len_swa_limit(
sliding_window_size: int, chunked_prefill_size: int, page_size: int sliding_window_size: int, chunked_prefill_size: int, page_size: int
) -> int: ) -> int:
# 1. a factor of 2x is because each prefill contains chunked_prefill_size tokens, # 1. a factor of 2x is because each prefill contains chunked_prefill_size tokens,
# and between prefills, we run swa_radix_cache.cache_unfinished_req(), # and between prefills, we run the tree cache's cache_unfinished_req(),
# so we unlock the previously locked nodes. # so we unlock the previously locked nodes.
# 2. max is to handle the case that chunked_prefill_size is larger than sliding_window_size. # 2. max is to handle the case that chunked_prefill_size is larger than sliding_window_size.
# in that case, each prefill contains chunked_prefill_size tokens, # in that case, each prefill contains chunked_prefill_size tokens,
@@ -2,7 +2,6 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Any, Callable, Dict from typing import TYPE_CHECKING, Any, Callable, Dict
from sglang.srt.mem_cache.swa_radix_cache import TreeNode as SWATreeNode
from sglang.srt.mem_cache.unified_radix_cache import UnifiedTreeNode from sglang.srt.mem_cache.unified_radix_cache import UnifiedTreeNode
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -18,8 +17,6 @@ def get_all_node_lock_refs(ctx: ScriptedContext) -> Dict[int, int]:
def _node_lock_ref(node: Any) -> int: def _node_lock_ref(node: Any) -> int:
if isinstance(node, SWATreeNode):
return node.full_lock_ref + node.swa_lock_ref
if isinstance(node, UnifiedTreeNode): if isinstance(node, UnifiedTreeNode):
return sum(cd.lock_ref for cd in node.component_data) return sum(cd.lock_ref for cd in node.component_data)
return node.lock_ref return node.lock_ref
@@ -1,7 +1,10 @@
"""DSV4 stress test for SWA radix cache + tombstone + retract interaction. """DSV4 stress test for SWA radix cache + tombstone + retract interaction.
Reproduces the assert in `swa_radix_cache.cache_unfinished_req`: Regression test for the former SWA `cache_unfinished_req` assertion:
assert old_prefix_len <= len(new_indices) assert old_prefix_len <= len(new_indices)
The unified cache reads `req.kv.cache_protected_len` and tolerates page_size - 1
of alignment slack, so this reproduces the historical trip conditions rather
than a line that still exists.
Trip conditions (all required): Trip conditions (all required):
1. Fork-only SWA leaf early-release on (`SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW=1`) 1. Fork-only SWA leaf early-release on (`SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW=1`)
@@ -88,7 +91,6 @@ class TestDSV4FlashSWARadixRetract(CustomTestCase):
env = { env = {
"SGLANG_DSV4_FP4_EXPERTS": "0", "SGLANG_DSV4_FP4_EXPERTS": "0",
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024",
"SGLANG_OPT_SWA_RADIX_CACHE_COMPACT": "0",
"SGLANG_TEST_RETRACT": "1", "SGLANG_TEST_RETRACT": "1",
"SGLANG_TEST_RETRACT_INTERVAL": "3", "SGLANG_TEST_RETRACT_INTERVAL": "3",
} }
@@ -130,7 +132,7 @@ class TestDSV4FlashSWARadixRetract(CustomTestCase):
"""Stress: 64 concurrent long-prompt reqs with long generation force """Stress: 64 concurrent long-prompt reqs with long generation force
retract under SWA pool pressure. Reqs share a 30k+ token prefix so retract under SWA pool pressure. Reqs share a 30k+ token prefix so
tombstoned leaves from retracted reqs are on the radix path of new tombstoned leaves from retracted reqs are on the radix path of new
reqs. Scheduler must not crash on the swa_radix_cache assert.""" reqs. Scheduler must not crash on the SWA insert assert."""
random.seed(0) random.seed(0)
concurrency = 64 concurrency = 64
@@ -9,7 +9,6 @@ import torch
from sglang.srt.kv_canary.radix_cache_walker import walk_radix_cache_for_canary from sglang.srt.kv_canary.radix_cache_walker import walk_radix_cache_for_canary
from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache, TreeNode
from sglang.srt.mem_cache.unified_cache.components import ( from sglang.srt.mem_cache.unified_cache.components import (
BASE_COMPONENT_TYPE, BASE_COMPONENT_TYPE,
ComponentType, ComponentType,
@@ -75,65 +74,6 @@ class TestSelfUnitRadixWalker(CustomTestCase):
result = walk_radix_cache_for_canary(radix_cache=cache, unlocked_only=True) result = walk_radix_cache_for_canary(radix_cache=cache, unlocked_only=True)
self.assertEqual(result.slot_indices.tolist(), [3, 4]) self.assertEqual(result.slot_indices.tolist(), [3, 4])
def test_walk_unlocked_only_uses_swa_full_lock_ref(self):
"""Verify SWA radix walking honors full-pool lock references."""
cache = SWARadixCache.__new__(SWARadixCache)
cache.device = self.device
cache.page_size = 1
cache.disable = False
root = TreeNode()
root.value = torch.tensor([], dtype=torch.int32, device=self.device)
cache.root_node = root
locked_child = TreeNode()
locked_child.value = torch.tensor([1, 2], dtype=torch.int32, device=self.device)
locked_child.parent = root
locked_child.full_lock_ref = 1
root.children[locked_child.id] = locked_child
unlocked_child = TreeNode()
unlocked_child.value = torch.tensor(
[3, 4], dtype=torch.int32, device=self.device
)
unlocked_child.parent = root
root.children[unlocked_child.id] = unlocked_child
result = walk_radix_cache_for_canary(radix_cache=cache, unlocked_only=True)
self.assertEqual(result.slot_indices.tolist(), [3, 4])
def test_swa_resident_only_skips_tombstoned_nodes(self):
"""Verify SWA radix walking skips nodes whose SWA storage was evicted."""
cache = SWARadixCache.__new__(SWARadixCache)
cache.device = self.device
cache.page_size = 1
cache.disable = False
root = TreeNode()
root.value = torch.tensor([], dtype=torch.int32, device=self.device)
cache.root_node = root
tombstoned_child = TreeNode()
tombstoned_child.value = torch.tensor(
[1, 2], dtype=torch.int32, device=self.device
)
tombstoned_child.parent = root
tombstoned_child.swa_tombstone = True
root.children[tombstoned_child.id] = tombstoned_child
resident_child = TreeNode()
resident_child.value = torch.tensor(
[3, 4], dtype=torch.int32, device=self.device
)
resident_child.parent = root
root.children[resident_child.id] = resident_child
result = walk_radix_cache_for_canary(
radix_cache=cache,
swa_resident_only=True,
)
self.assertEqual(result.slot_indices.tolist(), [3, 4])
def test_unified_swa_sweep_gates_on_swa_lock_not_full_lock(self): def test_unified_swa_sweep_gates_on_swa_lock_not_full_lock(self):
"""With unlocked_only + swa_resident_only, the sweep filters on the SWA """With unlocked_only + swa_resident_only, the sweep filters on the SWA
component lock: a FULL-locked node whose SWA lock was already released component lock: a FULL-locked node whose SWA lock was already released
@@ -7,24 +7,14 @@ import torch
from sglang.kernels.ops.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.kernels.ops.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
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.common import available_and_evictable_str
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.memory_pool import (
HybridLinearKVPool, HybridLinearKVPool,
HybridReqToTokenPool, HybridReqToTokenPool,
MambaPool, MambaPool,
) )
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import get_device from sglang.srt.utils import get_device
@@ -241,338 +231,10 @@ class TestMamba(unittest.TestCase):
view[0, 0, 0, 1, 0] = -1 view[0, 0, 0, 1, 0] = -1
self.assertEqual(view[0, 0, 1, 0, 0].item(), -1) self.assertEqual(view[0, 0, 1, 0, 0].item(), -1)
def test_mamba_radix_cache_1(self): def _setup_pools(self):
tree, allocator, req_to_token_pool, make_dummy_req = ( """Build the hybrid req/KV pools and an allocator for pool-level tests."""
self._setup_tree_and_allocator()
)
mamba_allocator = req_to_token_pool.mamba_allocator
mamba_pool = req_to_token_pool.mamba_pool
# test
print(
f"[Start] allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req1 = make_dummy_req()
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
assert len(req1_token_ids) == len(req1_kv_indices)
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(array("q", req1_token_ids))
result = tree.insert(
InsertParams(
key=key,
value=req1_kv_indices[: len(key)],
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req2 = make_dummy_req()
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
assert len(req2_token_ids) == len(req2_kv_indices)
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(array("q", req2_token_ids))
result = tree.insert(
InsertParams(
key=key,
value=req2_kv_indices[: len(key)],
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req3 = make_dummy_req()
req3_token_ids, req3_kv_indices = [10, 11, 12], allocator.alloc(3)
assert len(req3_token_ids) == len(req3_kv_indices)
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(array("q", req3_token_ids))
result = tree.insert(
InsertParams(
key=key,
value=req3_kv_indices[: len(key)],
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req4 = make_dummy_req()
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
assert len(req4_token_ids) == len(req4_kv_indices)
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(array("q", req4_token_ids))
result = tree.insert(
InsertParams(
key=key,
value=req4_kv_indices[: len(key)],
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
print(
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
tree.pretty_print()
full_num_tokens = 1
print(f"evicting {full_num_tokens} full token")
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")
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]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
assert len(kv_indices) == 0
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
assert len(kv_indices) == 7
assert len(last_node.key) == 2
req7_token_ids = [1, 2, 3, 4, 5, 6, 7]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req7_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req7: token_ids: {req7_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
assert len(kv_indices) == 7
assert len(last_node.key) == 2
mamba_num = 1
print(f"evicting {mamba_num} mamba")
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]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req8_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req8: token_ids: {req8_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
assert len(kv_indices) == 0
assert len(last_node.key) == 0
req9_token_ids = [1, 2, 3, 4, 5, 6, 7]
req9 = make_dummy_req()
result = tree.match_prefix(
MatchPrefixParams(
key=RadixKey(array("q", req9_token_ids)), req=req9, cow_mamba=True
)
)
kv_indices, last_node = result.device_indices, result.last_device_node
assert req9.kv.holds_mamba
assert torch.all(
mamba_pool.mamba_cache.conv[0][:, req9.kv.mamba_pool_idx]
== mamba_pool.mamba_cache.conv[0][:, last_node.mamba_value]
)
assert torch.all(
mamba_pool.mamba_cache.temporal[:, req9.kv.mamba_pool_idx]
== mamba_pool.mamba_cache.temporal[:, last_node.mamba_value]
)
print(tree.available_and_evictable_str())
print(available_and_evictable_str(tree))
tree.sanity_check()
def test_mamba_lru_match_refreshes_only_used_node(self):
"""A prefix-cache hit must refresh only the matched leaf's mamba state in
the mamba LRU, not its ancestors. Whole-chain refresh clustered a session's
states adjacently, so under mamba-pool pressure eviction dropped whole cold
sessions instead of the intermediate states reuse never needs. Guards against
reverting the mamba list to reset_node_and_parents_mru.
"""
tree, allocator, req_to_token_pool, make_dummy_req = (
self._setup_tree_and_allocator()
)
def insert(token_ids):
req = make_dummy_req()
kv = allocator.alloc(len(token_ids))
tree.insert(
InsertParams(
key=RadixKey(array("q", token_ids)),
value=kv,
mamba_value=req.kv.mamba_pool_idx.unsqueeze(0),
)
)
def match_leaf(token_ids):
return tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", token_ids)))
).last_device_node
def mamba_lru_mru_to_lru():
lst = tree.mamba_lru_list
order, x = [], getattr(lst.head, lst.nxt)
while x is not None and x is not lst.tail and x.id in lst.cache:
order.append(x)
x = getattr(x, lst.nxt)
return order
# Two independent sessions, each a 2-node mamba chain:
# root -> a1 -> b1 and root -> a2 -> b2
insert([1, 2, 3])
insert([1, 2, 3, 4, 5, 6])
insert([7, 8, 9])
insert([7, 8, 9, 10, 11, 12])
b1 = match_leaf([1, 2, 3, 4, 5, 6])
a1 = b1.parent
b2 = match_leaf([7, 8, 9, 10, 11, 12])
a2 = b2.parent
# Session 2 was matched last, so session 1's ancestor a1 is older than a2.
order = mamba_lru_mru_to_lru()
self.assertGreater(order.index(a1), order.index(a2))
# Re-access session 1. Only its consumed leaf (b1) moves to MRU; its ancestor
# a1 must stay put -- whole-chain reset would bump a1 right behind b1, making
# it newer than a2.
self.assertIs(match_leaf([1, 2, 3, 4, 5, 6]), b1)
order = mamba_lru_mru_to_lru()
self.assertIs(order[0], b1)
self.assertGreater(order.index(a1), order.index(a2))
tree.sanity_check()
def test_mamba_radix_cache_kv_events(self):
tree, allocator, _, make_dummy_req = self._setup_tree_and_allocator(
enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
stored_hashes = []
req1 = make_dummy_req()
key1 = RadixKey(array("q", [1, 2, 3]))
tree.insert(
InsertParams(
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
stored_events = [e for e in events if isinstance(e, BlockStored)]
self.assertEqual(len(stored_events), 1)
self.assertEqual(list(stored_events[0].token_ids), [1, 2, 3])
stored_hashes.extend(
block_hash for event in stored_events for block_hash in event.block_hashes
)
req2 = make_dummy_req()
key2 = RadixKey(array("q", [1, 2, 3, 4, 5]))
tree.insert(
InsertParams(
key=key2,
value=allocator.alloc(5)[: len(key2)],
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
stored_events = [e for e in events if isinstance(e, BlockStored)]
self.assertEqual(len(stored_events), 1)
self.assertEqual(list(stored_events[0].token_ids), [4, 5])
stored_hashes.extend(
block_hash for event in stored_events for block_hash in event.block_hashes
)
# Evicting an internal mamba state creates a tombstone but does not
# remove full-attention KV blocks, so it must not emit BlockRemoved.
result = tree.evict(EvictParams(num_tokens=0, mamba_num=1))
self.assertEqual(result.num_tokens_evicted, 0)
self.assertEqual(result.mamba_num_evicted, 1)
events = tree.take_events()
self.assertEqual([e for e in events if isinstance(e, BlockRemoved)], [])
result = tree.evict(EvictParams(num_tokens=1))
self.assertGreaterEqual(result.num_tokens_evicted, 1)
events = tree.take_events()
removed_hashes = _event_hashes(
[e for e in events if isinstance(e, BlockRemoved)]
)
self.assertCountEqual(removed_hashes, stored_hashes)
def test_mamba_radix_cache_kv_events_split_hash(self):
tree, allocator, _, make_dummy_req = self._setup_tree_and_allocator(
enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
req1 = make_dummy_req()
key1 = RadixKey(array("q", [1, 2, 3, 4]))
tree.insert(
InsertParams(
key=key1,
value=allocator.alloc(4)[: len(key1)],
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
first_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(first_insert_events), 1)
split_parent_hash = first_insert_events[0].block_hashes[1]
req2 = make_dummy_req()
key2 = RadixKey(array("q", [1, 2, 5, 6]))
tree.insert(
InsertParams(
key=key2,
value=allocator.alloc(4)[: len(key2)],
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
second_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 1)
self.assertEqual(list(second_insert_events[0].token_ids), [5, 6])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def _setup_tree_and_allocator(self, enable_kv_cache_events=False):
"""Helper to create a MambaRadixCache with allocator for testing."""
server_args = ServerArgs(model_path="dummy", page_size=1) server_args = ServerArgs(model_path="dummy", page_size=1)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise # The mamba pool reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model. # loads the HF config for self.model_path — impossible for the dummy model.
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE. # Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
server_args._mamba_cache_chunk_size = FLA_CHUNK_SIZE server_args._mamba_cache_chunk_size = FLA_CHUNK_SIZE
@@ -635,14 +297,6 @@ class TestMamba(unittest.TestCase):
kvcache=pool, kvcache=pool,
need_sort=False, need_sort=False,
) )
params = CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=1,
disable=False,
enable_kv_cache_events=enable_kv_cache_events,
)
tree = MambaRadixCache(params=params)
def make_dummy_req(): def make_dummy_req():
sampling_params = SamplingParams( sampling_params = SamplingParams(
@@ -658,7 +312,7 @@ class TestMamba(unittest.TestCase):
req_to_token_pool.alloc([req]) req_to_token_pool.alloc([req])
return req return req
return tree, allocator, req_to_token_pool, make_dummy_req return allocator, req_to_token_pool, make_dummy_req
# Qwen4-Exp's PLE N-gram window is 2 wide (ngram_size=3) and its "no history" # Qwen4-Exp's PLE N-gram window is 2 wide (ngram_size=3) and its "no history"
# sentinel is the eos id; pick a recognisable one for the tests. # sentinel is the eos id; pick a recognisable one for the tests.
@@ -700,7 +354,7 @@ class TestMamba(unittest.TestCase):
def test_slot_siblings_registered(self): def test_slot_siblings_registered(self):
"""Enabled PLE side states register on the pool that owns the slots; """Enabled PLE side states register on the pool that owns the slots;
disabled ones stay off so the host-offload payload keeps its legacy shape.""" disabled ones stay off so the host-offload payload keeps its legacy shape."""
_, _, base_pool, _ = self._setup_tree_and_allocator() _, base_pool, _ = self._setup_pools()
# The default hybrid setup has no PLE config: no siblings ride along. # The default hybrid setup has no PLE config: no siblings ride along.
self.assertEqual(len(base_pool.mamba_pool._slot_siblings), 0) self.assertEqual(len(base_pool.mamba_pool._slot_siblings), 0)
pool = self._setup_pool_with_ngram() pool = self._setup_pool_with_ngram()
@@ -813,7 +467,7 @@ class TestMamba(unittest.TestCase):
def test_mamba_pool_cpu_offload(self): def test_mamba_pool_cpu_offload(self):
"""MambaPool.get_cpu_copy / load_cpu_copy round-trips conv and temporal state.""" """MambaPool.get_cpu_copy / load_cpu_copy round-trips conv and temporal state."""
_, _, req_to_token_pool, _ = self._setup_tree_and_allocator() _, req_to_token_pool, _ = self._setup_pools()
mamba_pool = req_to_token_pool.mamba_pool mamba_pool = req_to_token_pool.mamba_pool
n = 3 n = 3
indices = req_to_token_pool.mamba_allocator.alloc(n) indices = req_to_token_pool.mamba_allocator.alloc(n)
@@ -862,7 +516,7 @@ class TestMamba(unittest.TestCase):
def test_hybrid_kv_pool_cpu_offload(self): def test_hybrid_kv_pool_cpu_offload(self):
"""HybridLinearKVPool.get_cpu_copy / load_cpu_copy saves and restores both """HybridLinearKVPool.get_cpu_copy / load_cpu_copy saves and restores both
the full-attention KV cache and Mamba state in a single round-trip.""" the full-attention KV cache and Mamba state in a single round-trip."""
_, allocator, req_to_token_pool, _ = self._setup_tree_and_allocator() allocator, req_to_token_pool, _ = self._setup_pools()
mamba_pool = req_to_token_pool.mamba_pool mamba_pool = req_to_token_pool.mamba_pool
hybrid_pool = allocator._kvcache # HybridLinearKVPool hybrid_pool = allocator._kvcache # HybridLinearKVPool
@@ -934,81 +588,6 @@ class TestMamba(unittest.TestCase):
mamba_cpu_none, "mamba_cpu should be None when mamba_indices=None" mamba_cpu_none, "mamba_cpu should be None when mamba_indices=None"
) )
def test_insert_prev_prefix_len(self):
"""Test that prev_prefix_len correctly controls which KV indices are freed
during insert, covering: full free, partial free across multi-node, and no free.
"""
tree, allocator, req_to_token_pool, make_dummy_req = (
self._setup_tree_and_allocator()
)
initial_avail = allocator.available_size()
# Step 1: Insert [1,2,3] to create first node
req1 = make_dummy_req()
key1 = RadixKey(array("q", [1, 2, 3]))
tree.insert(
InsertParams(
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
assert allocator.available_size() == initial_avail - 3
# 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(array("q", [1, 2, 3, 4, 5, 6, 7]))
result = tree.insert(
InsertParams(
key=key2,
value=allocator.alloc(7)[: len(key2)],
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=0,
)
)
assert result.prefix_len == 3
# alloc 7, freed 3 (dup prefix [0..2]), stored 4 in new node => net -4
assert allocator.available_size() == initial_avail - 3 - 4
avail_after_step2 = allocator.available_size()
# Step 3: Insert [1,2,3,4,5,6,7,8] with prev_prefix_len=2
# 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(array("q", [1, 2, 3, 4, 5, 6, 7, 8]))
result = tree.insert(
InsertParams(
key=key3,
value=allocator.alloc(8)[: len(key3)],
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=2,
)
)
assert result.prefix_len == 7
# alloc 8, freed 5, stored 1 => net -3
assert allocator.available_size() == avail_after_step2 - 3
avail_after_step3 = allocator.available_size()
# 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(array("q", [1, 2, 3, 4, 5, 6, 7, 8, 9]))
result = tree.insert(
InsertParams(
key=key4,
value=allocator.alloc(9)[: len(key4)],
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=8,
)
)
assert result.prefix_len == 8
# alloc 9, freed 0, stored 1 => net -9
assert allocator.available_size() == avail_after_step3 - 9
tree.sanity_check()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -45,7 +45,6 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams, MatchPrefixParams,
) )
from sglang.srt.mem_cache.events import KVCacheEventRecorder from sglang.srt.mem_cache.events import KVCacheEventRecorder
from sglang.srt.mem_cache.mamba_radix_cache import TreeNode as MambaTreeNode
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.utils import get_device from sglang.srt.utils import get_device
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -338,7 +337,7 @@ class TestTreeNode(unittest.TestCase):
def test_get_prefix_hash_values_not_shared_across_calls(self): def test_get_prefix_hash_values_not_shared_across_calls(self):
"""Regression guard for cached mutable prefix hash lists.""" """Regression guard for cached mutable prefix hash lists."""
for node_cls in (TreeNode, MambaTreeNode): for node_cls in (TreeNode,):
with self.subTest(node_cls=node_cls.__module__): with self.subTest(node_cls=node_cls.__module__):
root = node_cls() root = node_cls()
n1 = node_cls() n1 = node_cls()
@@ -1,550 +0,0 @@
"""Unit tests for SWA eviction boundary fixes.
Bug: when page_size > sliding_window_size, _evict_swa could advance the
eviction frontier to exactly page_floor(seq_len), making all tokens being
inserted into the radix tree fully evicted (case 3). _insert_helper had no
handling for this, creating an incorrect non-tombstone node that caused
inflated swa_evictable_size_, negative usage, and potential double-free.
Two-sided fix:
1. _evict_swa subtracts max(window, page) on the radix path (preventive).
2. _insert_helper early-returns on case 3 (defensive).
Tests use real tree/allocator/pool with mock Req/ScheduleBatch wrappers.
"""
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.managers.schedule_batch import ReqKvInfo, ScheduleBatch
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import free_swa_out_of_window_slots
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=13, stage="base-b", runner_config="1-gpu-large")
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd")
# ---------------------------------------------------------------------------
# Infrastructure helpers (shared setup, not logic)
# ---------------------------------------------------------------------------
def _swa_alloc(allocator, need_size):
"""Allocate from SWA allocator for any page_size.
SWATokenToKVPoolAllocator.alloc() asserts page_size == 1. For page_size > 1,
allocate from the underlying paged allocators directly and set up the mapping.
"""
if allocator.page_size == 1:
return allocator.alloc(need_size)
if need_size > allocator.full_attn_allocator.available_size():
return None
if need_size > allocator.swa_attn_allocator.available_size():
return None
full_indices = allocator.full_attn_allocator.alloc(need_size)
swa_indices = allocator.swa_attn_allocator.alloc(need_size)
assert full_indices is not None and swa_indices is not None
allocator.full_to_swa_index_mapping[full_indices] = swa_indices
return full_indices
def _build_swa_tree(page_size, sliding_window_size, kv_size=1024, kv_size_swa=512):
head_num, head_dim, num_layers, global_interval = 8, 128, 24, 4
dtype = torch.bfloat16
device = get_device()
full_ids = list(range(0, num_layers, global_interval))
swa_ids = [i for i in range(num_layers) if i not in set(full_ids)]
pool = ReqToTokenPool(
size=8, max_context_len=2048, device=device, enable_memory_saver=False
)
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_ids,
full_attention_layer_ids=full_ids,
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 = SWARadixCache(
params=CacheInitParams(
req_to_token_pool=pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
is_eagle=False,
sliding_window_size=sliding_window_size,
),
)
return tree, allocator, pool
def _make_req(req_pool_idx, token_ids, cache_protected_len, tree):
"""Mock Req with fields needed by _evict_swa and cache_finished_req."""
req = SimpleNamespace(
origin_input_ids=token_ids,
output_ids=[],
kv=ReqKvInfo(
req_pool_idx=req_pool_idx, cache_protected_len=cache_protected_len
),
extra_key=None,
cache_salt=None,
last_node=tree.root_node,
lock_receipt=DecLockRefParams(),
swa_prefix_lock_released=False,
prefix_indices=torch.tensor([], dtype=torch.int64, device=tree.device),
_kv_committed_len=len(token_ids),
)
return req
def _make_batch(tree, allocator, pool):
"""Mock ScheduleBatch with fields needed by _evict_swa."""
return SimpleNamespace(
tree_cache=tree,
req_to_token_pool=pool,
token_to_kv_pool_allocator=allocator,
)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestSWAEvictionBoundary(unittest.TestCase):
# -- Eviction formula: page_size > window --
def test_formula_page_gt_window_sweep(self):
"""Sweep page_size > window combinations. The -page_size fix must
prevent eviction from reaching page_floor(seq_len)."""
for page_size in [4, 8, 16, 32, 64, 128, 256]:
for window in [1, 2, 4, 8]:
if page_size <= window:
continue
tree, allocator, pool = _build_swa_tree(
page_size=page_size,
sliding_window_size=window,
kv_size=max(4096, page_size * 20),
kv_size_swa=max(2048, page_size * 10),
)
for seq_len in range(page_size + 1, page_size * 5):
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
if kv is None:
break
pool.write((0, slice(0, alloc_size)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
insert_len = seq_len // page_size * page_size
self.assertLess(
req.kv.swa_evicted_seqlen,
insert_len,
f"page={page_size}, win={window}, seq={seq_len}",
)
allocator.free(kv)
# -- Eviction formula: page_size <= window --
def test_formula_page_leq_window(self):
"""page_size <= window: -page_size fix causes no regression."""
page_size, window = 4, 8
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
for seq_len in [13, 17, 25, 33]:
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
pool.write((0, slice(0, alloc_size)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
insert_len = seq_len // page_size * page_size
self.assertLess(req.kv.swa_evicted_seqlen, insert_len)
tree.cache_finished_req(
req, is_insert=True, owned_kv_len=req._kv_committed_len
)
tree.sanity_check()
# -- Retention floor: never free past the last state checkpoint --
def test_retain_floor_clamps_eviction(self):
"""A hybrid cache keeps SWA down to the last state checkpoint, not to the
window behind the tail, because that is where a prefix match lands. The
floor must clamp the frontier even though the tail has moved far past it."""
page_size, window = 8, 16
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
seq_len = 200
checkpoint = 96
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
free_swa_out_of_window_slots(
req,
seq_len - 1,
sliding_window_size=window,
page_size=page_size,
req_to_token_pool=batch.req_to_token_pool,
token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator,
retain_floor=checkpoint - window,
)
# Without the floor this would reach page_floor(199 - 16) = 176.
self.assertLessEqual(req.kv.swa_evicted_seqlen, checkpoint - window)
self.assertEqual(req.kv.swa_evicted_seqlen % page_size, 0)
def test_retain_floor_ignored_for_chunk_cache(self):
"""Chunk cache builds no tree, so a retained checkpoint could never be
matched. Holding it would cost SWA slots for nothing."""
page_size, window = 8, 16
seq_len = 200
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
free_swa_out_of_window_slots(
req,
seq_len - 1,
sliding_window_size=window,
page_size=page_size,
req_to_token_pool=batch.req_to_token_pool,
token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator,
is_chunk_cache=True,
retain_floor=16,
)
expected = (seq_len - 1 - window) // page_size * page_size
self.assertEqual(req.kv.swa_evicted_seqlen, expected)
def test_retain_floor_none_matches_old_behaviour(self):
"""retain_floor=None must reproduce the pre-change frontier exactly, so a
cache without a second state stream is unaffected."""
page_size, window = 8, 16
seq_len = 200
frontiers = []
for floor in (None, "absent"):
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
kwargs = {} if floor == "absent" else {"retain_floor": None}
free_swa_out_of_window_slots(
req,
seq_len - 1,
sliding_window_size=window,
page_size=page_size,
req_to_token_pool=batch.req_to_token_pool,
token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator,
**kwargs,
)
frontiers.append(req.kv.swa_evicted_seqlen)
expected = (seq_len - 1 - max(window, page_size)) // page_size * page_size
self.assertEqual(frontiers[0], expected)
self.assertEqual(frontiers[1], expected)
def test_retain_floor_above_threshold_is_inert(self):
"""The floor is a min(), so a checkpoint that is already inside the window
must not hold anything extra."""
page_size, window = 8, 16
seq_len = 200
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
free_swa_out_of_window_slots(
req,
seq_len - 1,
sliding_window_size=window,
page_size=page_size,
req_to_token_pool=batch.req_to_token_pool,
token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator,
retain_floor=seq_len,
)
expected = (seq_len - 1 - max(window, page_size)) // page_size * page_size
self.assertEqual(req.kv.swa_evicted_seqlen, expected)
def test_retain_floor_does_not_unfree(self):
"""The frontier only advances. A floor arriving after slots were already
freed must not claim them back, which would double-free on the next pass."""
page_size, window = 8, 16
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
seq_len = 200
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
common_kwargs = dict(
sliding_window_size=window,
page_size=page_size,
req_to_token_pool=batch.req_to_token_pool,
token_to_kv_pool_allocator=batch.token_to_kv_pool_allocator,
)
free_swa_out_of_window_slots(req, seq_len - 1, **common_kwargs)
advanced = req.kv.swa_evicted_seqlen
self.assertGreater(advanced, 0)
free_swa_out_of_window_slots(req, seq_len - 1, retain_floor=0, **common_kwargs)
self.assertEqual(req.kv.swa_evicted_seqlen, advanced)
# -- Eviction formula: page_size == 1 --
def test_formula_page_size_1(self):
"""page_size=1: radix keeps max(window, page)=window, so the frontier is pre_len - window."""
page_size, window = 1, 4
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
for seq_len in range(window + 2, 30):
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
self.assertLess(req.kv.swa_evicted_seqlen, seq_len)
self.assertEqual(
req.kv.swa_evicted_seqlen, max(0, seq_len - 1 - max(window, page_size))
)
tree.cache_finished_req(
req, is_insert=True, owned_kv_len=req._kv_committed_len
)
tree.sanity_check()
# -- Eviction formula: no-op when seq too short --
def test_formula_noop_short_sequence(self):
"""pre_len - window - page_size < 0: eviction stays at 0."""
page_size, window = 8, 4
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
seq_len = page_size + window - 2 # = 10, formula gives 10-1-4-8 = -3
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
pool.write((0, slice(0, alloc_size)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
self.assertEqual(req.kv.swa_evicted_seqlen, 0)
# -- Insert case 1: swa_evicted <= total_prefix_length --
def test_insert_case1_evicted_within_matched(self):
"""Eviction within matched region. New tokens all non-tombstone."""
page_size, window = 8, 2
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
# First request: populate tree with 16 tokens (2 pages)
first_len = page_size * 2
kv1 = _swa_alloc(allocator, first_len)
pool.write((0, slice(0, first_len)), kv1)
req1 = _make_req(0, list(range(first_len)), 0, tree)
tree.cache_finished_req(
req1, is_insert=True, owned_kv_len=req1._kv_committed_len
)
tree.sanity_check()
# Second request: 24 tokens, first 16 overlap with tree
second_len = page_size * 3
kv2 = _swa_alloc(allocator, second_len)
pool.write((1, slice(0, second_len)), kv2)
req2 = _make_req(1, list(range(second_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
# pre_len=15: 15-2-8=5, floor to 8 -> 0. Eviction stays within matched.
ScheduleBatch._evict_swa(batch, req2, first_len - 1)
self.assertLessEqual(req2.kv.swa_evicted_seqlen, first_len)
swa_evictable_before = tree.swa_evictable_size_
tree.cache_finished_req(
req2, is_insert=True, owned_kv_len=req2._kv_committed_len
)
# New tokens [16, 24) should all be non-tombstone
new_tokens = second_len // page_size * page_size - first_len
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + new_tokens)
tree.sanity_check()
# -- Insert case 2: total_prefix_length < swa_evicted < total_length --
def test_insert_case2_partial_tombstone(self):
"""Partial eviction: some tombstone, some non-tombstone."""
page_size, window = 8, 2
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
# seq_len=25: insert_length=24, evicted should be 8 (1 page)
seq_len = page_size * 3 + 1
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
pool.write((0, slice(0, alloc_size)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
swa_evictable_before = tree.swa_evictable_size_
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
insert_len = seq_len // page_size * page_size
self.assertGreater(req.kv.swa_evicted_seqlen, 0, "Should have some eviction")
self.assertLess(req.kv.swa_evicted_seqlen, insert_len, "Should be partial")
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
non_tombstone = insert_len - req.kv.swa_evicted_seqlen
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + non_tombstone)
self.assertGreater(tree.full_evictable_size_, 0)
tree.sanity_check()
# -- Insert case 3: swa_evicted == total_length (defensive) --
def test_insert_case3_defensive_early_return(self):
"""Simulate OLD formula to trigger case 3. Defensive early return
must prevent non-tombstone node creation."""
page_size, window = 8, 2
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
# seq_len=11: OLD formula -> evicted=8 == insert_length=8
seq_len = page_size + window + 1
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
pool.write((0, slice(0, alloc_size)), kv)
# OLD formula (without -page_size)
pre_len = seq_len - 1
old_evicted = max(0, (pre_len - window) // page_size * page_size)
insert_len = seq_len // page_size * page_size
self.assertEqual(
old_evicted, insert_len, "Precondition: old formula hits boundary"
)
# Manually free SWA as _evict_swa would
allocator.free_swa(pool.req_to_token[0, :old_evicted])
req = _make_req(0, list(range(seq_len)), 0, tree)
req.kv.swa_evicted_seqlen = old_evicted
swa_evictable_before = tree.swa_evictable_size_
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before)
# -- Integration: multiple decode turns --
def test_multiple_decodes(self):
"""Multiple decode turns with page_size > window. No over-eviction,
tree stays consistent throughout."""
page_size, window = 8, 2
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
for turn in range(4):
seq_len = page_size * (turn + 2) + 1
idx = turn % pool.size
alloc_size = (seq_len + page_size - 1) // page_size * page_size
kv = _swa_alloc(allocator, alloc_size)
assert kv is not None, f"Alloc failed at turn {turn}"
pool.write((idx, slice(0, alloc_size)), kv)
req = _make_req(idx, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
insert_len = seq_len // page_size * page_size
self.assertLess(req.kv.swa_evicted_seqlen, insert_len, f"turn {turn}")
tree.cache_finished_req(
req, is_insert=True, owned_kv_len=req._kv_committed_len
)
tree.sanity_check()
# -- Integration: page_size=1 full flow --
def test_page_size_1_full_flow(self):
"""End-to-end with page_size=1. Fix is near no-op."""
page_size, window = 1, 4
tree, allocator, pool = _build_swa_tree(
page_size=page_size, sliding_window_size=window
)
for seq_len in [10, 20, 30]:
kv = _swa_alloc(allocator, seq_len)
pool.write((0, slice(0, seq_len)), kv)
req = _make_req(0, list(range(seq_len)), 0, tree)
batch = _make_batch(tree, allocator, pool)
ScheduleBatch._evict_swa(batch, req, seq_len - 1)
self.assertEqual(
req.kv.swa_evicted_seqlen, max(0, seq_len - 1 - max(window, page_size))
)
tree.cache_finished_req(
req, is_insert=True, owned_kv_len=req._kv_committed_len
)
tree.sanity_check()
if __name__ == "__main__":
unittest.main()
@@ -1,440 +0,0 @@
"""Regression for SWA lock release lifecycle.
Hybrid-SWA early-release protocol: once a request's decode position passes
the sliding window, drop its prefill SWA lock without touching the full
lock, freeing SWA pages back to LRU.
Covers:
- SWARadixCache.dec_swa_lock_only (leaf tombstone + free, internal protected->evictable)
- SWARadixCache.dec_lock_ref(skip_swa=True)
- SWARadixCache.evict swa branch for leaf with full_lock_ref > 0
- SWARadixCache._delete_leaf skipping swa_evictable_size_ on tombstoned leaves
"""
import unittest
from array import array
import torch
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
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.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=12, suite="stage-b-test-1-gpu-small-amd")
def _build_tree(
*,
sliding_window_size: int = 4,
page_size: int = 1,
kv_size: int = 128,
kv_size_swa: int = 64,
):
head_num, head_dim, num_layers, global_interval = 8, 128, 24, 4
dtype = torch.bfloat16
device = get_device()
full_ids = list(range(0, num_layers, global_interval))
swa_ids = [i for i in range(num_layers) if i not in set(full_ids)]
pool = ReqToTokenPool(
size=8, max_context_len=256, device=device, enable_memory_saver=False
)
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_ids,
full_attention_layer_ids=full_ids,
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 = SWARadixCache(
params=CacheInitParams(
req_to_token_pool=pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
is_eagle=False,
sliding_window_size=sliding_window_size,
),
)
return tree, allocator, pool
def _swa_alloc(allocator, need_size):
"""Allocate from SWA allocator for any page_size.
SWATokenToKVPoolAllocator.alloc() asserts page_size == 1; for page_size > 1
we drive the underlying paged allocators directly (mirrors the helper in
test_swa_eviction_boundary.py). Required: need_size is a multiple of
page_size when page_size > 1.
"""
if allocator.page_size == 1:
return allocator.alloc(need_size)
assert need_size % allocator.page_size == 0, (
f"page_size > 1 requires page-aligned alloc, got {need_size=} "
f"with {allocator.page_size=}"
)
if need_size > allocator.full_attn_allocator.available_size():
return None
if need_size > allocator.swa_attn_allocator.available_size():
return None
full_indices = allocator.full_attn_allocator.alloc(need_size)
swa_indices = allocator.swa_attn_allocator.alloc(need_size)
assert full_indices is not None and swa_indices is not None
allocator.full_to_swa_index_mapping[full_indices] = swa_indices
return full_indices
def _insert_chain(tree, allocator, token_ids):
token_ids = array("q", token_ids)
indices = _swa_alloc(allocator, len(token_ids))
assert indices is not None
tree.insert(InsertParams(key=RadixKey(token_ids), value=indices))
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(token_ids)))
return match.last_device_node
def _release_swa_lock_chain_in_place(tree, leaf, swa_uuid_for_lock):
# Mirrors dec_swa_lock_only's non-tombstone arm (protected->evictable on
# internal nodes) but skips the leaf-free + tombstone step, to construct
# the post-revival state where SWA was already early-released yet the
# leaf is back in swa_lru_list with full_lock_ref still > 0.
node = leaf
while node is not tree.root_node:
if node.swa_lock_ref > 0:
if node.swa_lock_ref == 1:
tree.swa_protected_size_ -= len(node.value)
tree.swa_evictable_size_ += len(node.value)
node.swa_lock_ref -= 1
if swa_uuid_for_lock and node.swa_uuid == swa_uuid_for_lock:
break
node = node.parent
class TestSWALockReleaseLifecycle(CustomTestCase):
"""Each test pins one component of the early-release fix; method names
are prefixed with the API surface they exercise so pytest output groups
them naturally."""
def test_dec_swa_lock_only_leaf_tombstones_and_frees(self):
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
self.assertEqual(len(leaf.value), 8)
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
self.assertIsNotNone(swa_uuid)
swa_avail_before = allocator.swa_available_size()
full_avail_before = allocator.full_available_size()
self.assertEqual(leaf.swa_lock_ref, 1)
self.assertEqual(leaf.full_lock_ref, 1)
self.assertFalse(leaf.swa_tombstone)
self.assertTrue(tree.swa_lru_list.in_list(leaf))
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
self.assertTrue(leaf.swa_tombstone)
self.assertFalse(tree.swa_lru_list.in_list(leaf))
self.assertEqual(leaf.swa_lock_ref, 0)
self.assertEqual(
allocator.swa_available_size(), swa_avail_before + len(leaf.value)
)
self.assertEqual(leaf.full_lock_ref, 1)
self.assertEqual(allocator.full_available_size(), full_avail_before)
# sanity_check forbids live locks; release the full half before checking.
tree.dec_lock_ref(
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
)
tree.sanity_check()
def test_dec_swa_lock_only_internal_no_tombstone_no_free(self):
# Two siblings force an internal node at the shared prefix.
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf_a = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
_insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 9])
# Post-split: leaf_a now carries [8] only, parent holds the shared 7.
self.assertEqual(len(leaf_a.value), 1)
internal = leaf_a.parent
self.assertGreater(len(internal.children), 1)
self.assertEqual(len(internal.value), 7)
inc_res = tree.inc_lock_ref(leaf_a)
swa_uuid = inc_res.swa_uuid_for_lock
# window=4, value 1 (leaf) + 7 (internal): swa lock chain ends at internal.
self.assertEqual(swa_uuid, internal.swa_uuid)
swa_protected_before = tree.swa_protected_size_
swa_evictable_before = tree.swa_evictable_size_
swa_avail_before = allocator.swa_available_size()
tree.dec_swa_lock_only(leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
self.assertFalse(internal.swa_tombstone)
self.assertTrue(tree.swa_lru_list.in_list(internal))
self.assertEqual(internal.swa_lock_ref, 0)
self.assertEqual(
tree.swa_protected_size_, swa_protected_before - (len(leaf_a.value) + 7)
)
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before + 7)
self.assertEqual(
allocator.swa_available_size(), swa_avail_before + len(leaf_a.value)
)
tree.dec_lock_ref(
leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
)
tree.sanity_check()
def test_dec_lock_ref_skip_swa_true_drops_full_only(self):
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
self.assertTrue(leaf.swa_tombstone)
self.assertEqual(leaf.full_lock_ref, 1)
swa_avail_after_release = allocator.swa_available_size()
swa_protected_after_release = tree.swa_protected_size_
# Without skip_swa, dec_lock_ref would assert on the swa_tombstone leaf.
tree.dec_lock_ref(
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
)
self.assertEqual(leaf.full_lock_ref, 0)
self.assertEqual(allocator.swa_available_size(), swa_avail_after_release)
self.assertEqual(tree.swa_protected_size_, swa_protected_after_release)
tree.sanity_check()
def test_dec_lock_ref_skip_swa_false_drops_both(self):
# Default skip_swa=False must keep legacy behavior intact.
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
full_avail_before = allocator.full_available_size()
swa_avail_before = allocator.swa_available_size()
tree.dec_lock_ref(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
self.assertEqual(leaf.full_lock_ref, 0)
self.assertEqual(leaf.swa_lock_ref, 0)
self.assertEqual(tree.full_protected_size_, 0)
self.assertEqual(tree.swa_protected_size_, 0)
# dec_lock_ref releases locks but doesn't free; eviction does.
self.assertEqual(allocator.full_available_size(), full_avail_before)
self.assertEqual(allocator.swa_available_size(), swa_avail_before)
tree.sanity_check()
def test_evict_swa_leaf_with_full_lock_tombstones_in_place(self):
# Large window so inc_lock_ref locks the entire SWA chain.
tree, allocator, _ = _build_tree(sliding_window_size=64)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4])
self.assertEqual(len(leaf.value), 4)
inc_res = tree.inc_lock_ref(leaf)
_release_swa_lock_chain_in_place(tree, leaf, inc_res.swa_uuid_for_lock)
self.assertEqual(leaf.full_lock_ref, 1)
self.assertEqual(leaf.swa_lock_ref, 0)
self.assertFalse(leaf.swa_tombstone)
self.assertTrue(tree.swa_lru_list.in_list(leaf))
swa_avail_before = allocator.swa_available_size()
swa_evictable_before = tree.swa_evictable_size_
# num_tokens=0 skips the full eviction loop; swa loop hits the new branch.
evict_res = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=4))
self.assertGreaterEqual(evict_res.swa_num_tokens_evicted, 4)
self.assertTrue(leaf.swa_tombstone)
self.assertFalse(tree.swa_lru_list.in_list(leaf))
self.assertEqual(leaf.full_lock_ref, 1)
self.assertEqual(
allocator.swa_available_size(), swa_avail_before + len(leaf.value)
)
# Full lock prevents _delete_leaf, so the node stays attached.
self.assertIs(leaf.parent.children[leaf.key.child_key(tree.page_size)], leaf)
self.assertEqual(
tree.swa_evictable_size_, swa_evictable_before - len(leaf.value)
)
tree.dec_lock_ref(
leaf,
DecLockRefParams(swa_uuid_for_lock=inc_res.swa_uuid_for_lock),
skip_swa=True,
)
tree.sanity_check()
def test_delete_leaf_skips_swa_size_on_tombstone(self):
# Tombstone removes the count once; _delete_leaf must not subtract again.
tree, allocator, _ = _build_tree(sliding_window_size=4)
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
self.assertTrue(leaf.swa_tombstone)
swa_evictable_before_delete = tree.swa_evictable_size_
tree.full_lru_list.remove_node(leaf)
tree._delete_leaf(leaf)
self.assertEqual(tree.swa_evictable_size_, swa_evictable_before_delete)
def test_dec_swa_lock_only_leaf_page_size_variants(self):
"""Single-leaf tombstone+free across all (page_size, window) regimes.
Sweep covers:
- window multiple of page_size (page_size=2, window=4)
- page_size > window (page_size=8, window=4)
- window not multiple of page (page_size=4, window=6)
With page_size > 1, _swa_alloc routes through the paged allocators;
free_swa(leaf.value) must release exactly len(leaf.value) tokens
(page-aligned) regardless of how page_size relates to the window.
"""
for page_size, window in [(2, 4), (8, 4), (4, 6)]:
with self.subTest(page_size=page_size, window=window):
tree, allocator, _ = _build_tree(
sliding_window_size=window,
page_size=page_size,
kv_size=max(128, 32 * page_size),
kv_size_swa=max(64, 16 * page_size),
)
n_tokens = max(window, 2 * page_size)
n_tokens = (n_tokens + page_size - 1) // page_size * page_size
leaf = _insert_chain(tree, allocator, list(range(1, n_tokens + 1)))
self.assertEqual(len(leaf.value), n_tokens)
self.assertEqual(len(leaf.value) % page_size, 0)
inc_res = tree.inc_lock_ref(leaf)
swa_uuid = inc_res.swa_uuid_for_lock
self.assertIsNotNone(
swa_uuid,
f"inc_lock_ref must reach the window with leaf.value="
f"{len(leaf.value)} >= window={window}",
)
swa_avail_before = allocator.swa_available_size()
full_avail_before = allocator.full_available_size()
tree.dec_swa_lock_only(
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid)
)
self.assertTrue(leaf.swa_tombstone)
self.assertFalse(tree.swa_lru_list.in_list(leaf))
self.assertEqual(leaf.swa_lock_ref, 0)
self.assertEqual(
allocator.swa_available_size(),
swa_avail_before + len(leaf.value),
"free_swa must release the leaf's full page-aligned slot count",
)
self.assertEqual(leaf.full_lock_ref, 1)
self.assertEqual(allocator.full_available_size(), full_avail_before)
tree.dec_lock_ref(
leaf,
DecLockRefParams(swa_uuid_for_lock=swa_uuid),
skip_swa=True,
)
tree.sanity_check()
def test_dec_swa_lock_only_internal_page_size_gt_1(self):
"""Internal-node chain release with page_size > 1.
Two siblings sharing a page-aligned prefix force a radix split on a
page boundary. The swa lock chain therefore spans leaf -> internal,
and dec_swa_lock_only must:
- tombstone the leaf and free len(leaf.value) SWA tokens
- flip the internal node from protected -> evictable (no free,
no tombstone)
"""
page_size, window = 2, 6
tree, allocator, _ = _build_tree(
sliding_window_size=window, page_size=page_size
)
# Shared prefix len 4 (2 pages); divergent suffix len 2 (1 page each).
leaf_a = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6])
_insert_chain(tree, allocator, [1, 2, 3, 4, 7, 8])
self.assertEqual(len(leaf_a.value), 2)
internal = leaf_a.parent
self.assertGreater(len(internal.children), 1)
self.assertEqual(len(internal.value), 4)
inc_res = tree.inc_lock_ref(leaf_a)
swa_uuid = inc_res.swa_uuid_for_lock
# leaf_a (2) + internal (4) = 6 >= window=6, so uuid stops at internal.
self.assertEqual(swa_uuid, internal.swa_uuid)
swa_protected_before = tree.swa_protected_size_
swa_evictable_before = tree.swa_evictable_size_
swa_avail_before = allocator.swa_available_size()
tree.dec_swa_lock_only(leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
# Leaf side: tombstoned and pages freed.
self.assertTrue(leaf_a.swa_tombstone)
self.assertFalse(tree.swa_lru_list.in_list(leaf_a))
self.assertEqual(
allocator.swa_available_size(),
swa_avail_before + len(leaf_a.value),
)
# Internal side: protected -> evictable, still in lru, no free.
self.assertFalse(internal.swa_tombstone)
self.assertTrue(tree.swa_lru_list.in_list(internal))
self.assertEqual(internal.swa_lock_ref, 0)
self.assertEqual(
tree.swa_protected_size_,
swa_protected_before - (len(leaf_a.value) + len(internal.value)),
)
self.assertEqual(
tree.swa_evictable_size_,
swa_evictable_before + len(internal.value),
)
tree.dec_lock_ref(
leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
)
tree.sanity_check()
if __name__ == "__main__":
unittest.main()
@@ -1,12 +1,10 @@
import unittest import unittest
from array import array
from types import SimpleNamespace from types import SimpleNamespace
from unittest import mock from unittest import mock
from unittest.mock import patch from unittest.mock import patch
import torch import torch
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.environ import InvariantCheckLevel, envs from sglang.srt.environ import InvariantCheckLevel, envs
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.swa import ( from sglang.srt.mem_cache.allocator.swa import (
@@ -15,21 +13,12 @@ from sglang.srt.mem_cache.allocator.swa import (
) )
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache, BasePrefixCache,
DecLockRefParams,
EvictParams,
EvictResult,
InsertParams,
MatchPrefixParams,
) )
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
available_and_evictable_str,
free_kv_row_segments, free_kv_row_segments,
) )
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.utils import get_device from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -38,10 +27,6 @@ register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-large")
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd") register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd")
def _event_hashes(events):
return [block_hash for event in events for block_hash in event.block_hashes]
class _DummyReq: class _DummyReq:
def __init__(self): def __init__(self):
self._kv_committed_len = 0 self._kv_committed_len = 0
@@ -100,18 +85,7 @@ def _build_swa_tree(
need_sort=False, need_sort=False,
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
) )
tree = SWARadixCache( return allocator, req_to_token_pool
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
is_eagle=is_eagle,
sliding_window_size=sliding_window_size,
enable_kv_cache_events=enable_kv_cache_events,
),
)
return tree, allocator, req_to_token_pool
def _sync_error(fn): def _sync_error(fn):
@@ -164,23 +138,6 @@ def _swa_alloc(allocator, need_size):
return full_indices return full_indices
def _insert(tree, allocator, token_ids):
indices = _swa_alloc(allocator, len(token_ids))
assert indices is not None
tree.insert(InsertParams(key=RadixKey(array("q", token_ids)), value=indices))
def _insert_chain(tree, allocator, token_ids):
_insert(tree, allocator, token_ids)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", token_ids))))
return match.last_device_node
def _expected_tail_size(window: int, page_size: int) -> int:
"""Mirror of _maybe_split_leaf_for_swa_lock's tail_size formula."""
return (window + page_size - 1) // page_size * page_size
class TestSWA(unittest.TestCase): class TestSWA(unittest.TestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -190,71 +147,9 @@ class TestSWA(unittest.TestCase):
def tearDownClass(cls): def tearDownClass(cls):
pass pass
def test_swa_radix_cache_kv_events(self):
tree, allocator, _ = _build_swa_tree(
is_eagle=False, enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
_insert(tree, allocator, [1, 2, 3, 4])
first_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(first_insert_events), 1)
self.assertEqual(list(first_insert_events[0].token_ids), [1, 2, 3, 4])
_insert(tree, allocator, [1, 2, 3, 4, 5, 6])
second_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 1)
self.assertEqual(list(second_insert_events[0].token_ids), [5, 6])
stored_hashes = [
block_hash
for event in first_insert_events + second_insert_events
for block_hash in event.block_hashes
]
# Evicting only SWA tokens tombstones nodes but keeps full KV blocks.
result = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=1))
self.assertEqual(result.num_tokens_evicted, 0)
self.assertGreaterEqual(result.swa_num_tokens_evicted, 1)
self.assertEqual(
[e for e in tree.take_events() if isinstance(e, BlockRemoved)], []
)
result = tree.evict(EvictParams(num_tokens=1, swa_num_tokens=0))
self.assertGreaterEqual(result.num_tokens_evicted, 1)
removed_hashes = _event_hashes(
[e for e in tree.take_events() if isinstance(e, BlockRemoved)]
)
self.assertCountEqual(removed_hashes, stored_hashes)
def test_swa_radix_cache_kv_events_split_hash(self):
tree, allocator, _ = _build_swa_tree(
is_eagle=False, enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
_insert(tree, allocator, [1, 2, 3, 4])
first_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(first_insert_events), 1)
split_parent_hash = first_insert_events[0].block_hashes[1]
_insert(tree, allocator, [1, 2, 5, 6])
second_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 1)
self.assertEqual(list(second_insert_events[0].token_ids), [5, 6])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def test_swa_memory_pool_paged_free_clears_full_page_mapping(self): def test_swa_memory_pool_paged_free_clears_full_page_mapping(self):
page_size = 4 page_size = 4
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=page_size, page_size=page_size,
kv_size=16, kv_size=16,
@@ -281,7 +176,7 @@ class TestSWA(unittest.TestCase):
"""Clearing the full-to-SWA mapping must not block the stream; writing a """Clearing the full-to-SWA mapping must not block the stream; writing a
host-resident scalar into it does. host-resident scalar into it does.
""" """
_, allocator, _ = _build_swa_tree(is_eagle=False) allocator, _ = _build_swa_tree(is_eagle=False)
full_indices = _swa_alloc(allocator, 4) full_indices = _swa_alloc(allocator, 4)
mapping = allocator.full_to_swa_index_mapping mapping = allocator.full_to_swa_index_mapping
@@ -307,7 +202,7 @@ class TestSWA(unittest.TestCase):
self._free_swa_group_owns_deferred_indices(page_size) self._free_swa_group_owns_deferred_indices(page_size)
def _free_swa_group_owns_deferred_indices(self, page_size): def _free_swa_group_owns_deferred_indices(self, page_size):
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=page_size, page_size=page_size,
kv_size=32 * page_size, kv_size=32 * page_size,
@@ -344,7 +239,7 @@ class TestSWA(unittest.TestCase):
) )
def test_free_swa_group_owns_mapping_at_enqueue_time(self): def test_free_swa_group_owns_mapping_at_enqueue_time(self):
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
kv_size=8, kv_size=8,
kv_size_swa=8, kv_size_swa=8,
@@ -376,7 +271,7 @@ class TestSWA(unittest.TestCase):
) )
def _build_two_mapped_slots(self, page_size=1): def _build_two_mapped_slots(self, page_size=1):
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=page_size, page_size=page_size,
kv_size=8 * page_size, kv_size=8 * page_size,
@@ -447,524 +342,6 @@ class TestSWA(unittest.TestCase):
allocator.full_to_swa_index_mapping[indices], indices allocator.full_to_swa_index_mapping[indices], indices
) )
def test_swa_radix_cache_1(self):
# args
req_size = 10
max_context_len = 128
kv_size = 128
kv_size_swa = 64
page_size = 1
sliding_window_size = 4
head_num = 8
head_dim = 128
num_layers = 48
global_interval = 4
dtype = torch.bfloat16
device = get_device()
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
full_attention_layer_ids_set = set(full_attention_layer_ids)
swa_attention_layer_ids = [
i for i in range(num_layers) if i not in full_attention_layer_ids_set
]
# setup req to token pool
req_to_token_pool = ReqToTokenPool(
size=req_size,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
# setup kv pool
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,
device=device,
)
# setup token to kv pool allocator
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size_swa,
page_size=page_size,
dtype=dtype,
device=device,
kvcache=kv_pool,
need_sort=False,
)
# setup radix cache
tree = SWARadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
disable=False,
page_size=page_size,
sliding_window_size=sliding_window_size,
),
)
# test
print(
f"[Start] allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
self.assertEqual(len(req1_token_ids), len(req1_kv_indices))
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(array("q", 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()}"
)
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
self.assertEqual(len(req2_token_ids), len(req2_kv_indices))
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(array("q", 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()}"
)
req3_token_ids, req3_kv_indices = [10, 11, 12], allocator.alloc(3)
self.assertEqual(len(req3_token_ids), len(req3_kv_indices))
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(array("q", 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()}"
)
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
self.assertEqual(len(req4_token_ids), len(req4_kv_indices))
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(array("q", 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()}"
)
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(
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(
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(
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
)
tree.pretty_print()
req5_token_ids = [1, 2, 3, 4, 5]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
self.assertEqual(len(kv_indices), 0)
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
self.assertEqual(len(kv_indices), 7)
self.assertEqual(len(last_node.key), 2)
self.assertEqual(last_node.key.token_ids[0], 60)
self.assertEqual(last_node.key.token_ids[1], 70)
print(tree.available_and_evictable_str())
print(available_and_evictable_str(tree))
tree.sanity_check()
def test_swa_radix_cache_eagle(self):
# args
req_size = 10
max_context_len = 128
kv_size = 128
kv_size_swa = 64
page_size = 1
sliding_window_size = 4
head_num = 8
head_dim = 128
num_layers = 48
global_interval = 4
dtype = torch.bfloat16
device = get_device()
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
full_attention_layer_ids_set = set(full_attention_layer_ids)
swa_attention_layer_ids = [
i for i in range(num_layers) if i not in full_attention_layer_ids_set
]
# setup req to token pool
req_to_token_pool = ReqToTokenPool(
size=req_size,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
# setup kv pool
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,
device=device,
)
# setup token to kv pool allocator
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size_swa,
page_size=page_size,
dtype=dtype,
device=device,
kvcache=kv_pool,
need_sort=False,
)
# setup radix cache
tree = SWARadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
is_eagle=True,
sliding_window_size=sliding_window_size,
),
)
# test
print(
f"[Start] allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
self.assertEqual(len(req1_token_ids), len(req1_kv_indices))
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(array("q", 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(
f"req1: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
self.assertEqual(len(req2_token_ids), len(req2_kv_indices))
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(array("q", 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(
f"req2: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
req3_token_ids, req3_kv_indices = [10, 11, 12], allocator.alloc(3)
self.assertEqual(len(req3_token_ids), len(req3_kv_indices))
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(array("q", 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(
f"req3: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
self.assertEqual(len(req4_token_ids), len(req4_kv_indices))
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(array("q", 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(
f"req4: prefix_len: {prefix_len}, allocator swa available size: {allocator.swa_available_size()}, full available size: {allocator.full_available_size()}"
)
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")
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")
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")
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]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
self.assertEqual(len(kv_indices), 0) # no swa prefix matched
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
)
self.assertEqual(len(kv_indices), 6)
self.assertEqual(len(last_node.key), 2)
# Bigram view: token_ids holds raw tokens; iteration yields bigram tuples.
self.assertTrue(last_node.key.is_bigram)
self.assertEqual(list(last_node.key), [(5, 60), (60, 70)])
def test_swa_cache_finished_req_eagle_uses_cache_protected_len_and_bigram_key(self):
tree, allocator, req_to_token_pool = _build_swa_tree(is_eagle=True)
# Case 1: is_insert=True should pass bigram key and use cache_protected_len.
req = _DummyReq()
req.kv.req_pool_idx = 0
req.origin_input_ids = array("q", [1, 2, 3, 4, 5, 6])
req.output_ids = array("q")
req._kv_committed_len = len(req.origin_input_ids)
kv_indices = allocator.alloc(req._kv_committed_len)
req_to_token_pool.write(
(req.kv.req_pool_idx, slice(0, req._kv_committed_len)), kv_indices
)
req.extra_key = None
req.cache_salt = None
req.last_node = tree.root_node
req.lock_receipt = DecLockRefParams()
req.kv.swa_evicted_seqlen = 0
req.kv.cache_protected_len = 1
# Intentionally mismatch to ensure code does not use len(prefix_indices).
req.prefix_indices = torch.tensor([7, 8, 9, 10, 11], device=tree.device)
captured = {}
original_insert = tree.insert
def wrapped_insert(params):
captured["prev_prefix_len"] = params.prev_prefix_len
captured["is_bigram"] = params.key.is_bigram
captured["key_len"] = len(params.key)
return original_insert(params)
tree.insert = wrapped_insert
tree.cache_finished_req(req, is_insert=True, owned_kv_len=req._kv_committed_len)
self.assertEqual(captured["prev_prefix_len"], req.kv.cache_protected_len)
self.assertTrue(captured["is_bigram"])
self.assertEqual(captured["key_len"], len(req.origin_input_ids) - 1)
# Case 2: is_insert=False should free [cache_protected_len:page_aligned_len]
# even when len(prefix_indices) is intentionally larger.
req2 = _DummyReq()
req2.kv.req_pool_idx = 1
req2.origin_input_ids = array("q", [11, 12, 13, 14, 15, 16])
req2.output_ids = array("q")
req2._kv_committed_len = len(req2.origin_input_ids)
kv_indices2 = allocator.alloc(req2._kv_committed_len)
req_to_token_pool.write(
(req2.kv.req_pool_idx, slice(0, req2._kv_committed_len)), kv_indices2
)
req2.extra_key = None
req2.cache_salt = None
req2.last_node = tree.root_node
req2.lock_receipt = DecLockRefParams()
req2.kv.swa_evicted_seqlen = 0
req2.kv.cache_protected_len = 1
req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device)
freed_lens = []
original_free_segment = allocator.free_segment
def wrapped_free_segment(indices, *, start_pos):
freed_lens.append(int(indices.numel()))
return original_free_segment(indices, start_pos=start_pos)
allocator.free_segment = wrapped_free_segment
tree.cache_finished_req(
req2, is_insert=False, owned_kv_len=req2._kv_committed_len
)
# EAGLE + page_size=1 => page_aligned_len = committed_len - 1 = 5
# Expected frees:
# overlap range [1:5] -> 4
# tail range [5:] -> 1
self.assertEqual(freed_lens, [4, 1])
# Optimization: SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.
# Splits a freshly-inserted leaf at the (page-aligned) sliding-window
# boundary so a future inc_lock_ref protects only ~sliding_window_size SWA
# tokens instead of the whole chunked-prefill chain.
class TestSWASplitLeafOnInsert(CustomTestCase):
def _insert_and_lock(self, *, window, page_size, leaf_len, flag_on):
tree, allocator, _ = _build_swa_tree(
is_eagle=False,
kv_size=128,
kv_size_swa=64,
sliding_window_size=window,
page_size=page_size,
)
token_ids = list(range(leaf_len))
with envs.SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.override(flag_on):
leaf = _insert_chain(tree, allocator, token_ids)
result = tree.inc_lock_ref(leaf)
return tree, leaf, result
def test_flag_off_protects_full_leaf(self):
tree, leaf, _ = self._insert_and_lock(
window=4, page_size=1, leaf_len=12, flag_on=False
)
self.assertEqual(len(leaf.value), 12)
self.assertEqual(tree.swa_protected_size_, 12)
def test_flag_on_caps_protection_at_window(self):
# (window, page_size, leaf_len, expected_tail_size); leaf_len picked
# > tail_size and page-aligned for page_size > 1.
cases = [
(4, 1, 12, 4),
(4, 1, 5, 4),
(1, 1, 5, 1),
(4, 2, 12, 4),
(8, 2, 12, 8),
(4, 4, 12, 4),
# window NOT page-aligned -> tail rounds up to page boundary.
(3, 2, 12, 4),
(5, 4, 12, 8),
(3, 4, 12, 4),
]
for window, page_size, leaf_len, expected_tail in cases:
with self.subTest(window=window, page_size=page_size, leaf_len=leaf_len):
self.assertEqual(_expected_tail_size(window, page_size), expected_tail)
tree, leaf, _ = self._insert_and_lock(
window=window,
page_size=page_size,
leaf_len=leaf_len,
flag_on=True,
)
self.assertEqual(len(leaf.value), expected_tail)
self.assertEqual(tree.swa_protected_size_, expected_tail)
def test_flag_on_no_split_when_leaf_within_window(self):
# leaf_len <= tail_size: split must no-op.
cases = [
(4, 1, 4),
(4, 1, 3),
(4, 2, 4),
(3, 2, 4),
(8, 2, 4),
(4, 4, 4),
]
for window, page_size, leaf_len in cases:
with self.subTest(window=window, page_size=page_size, leaf_len=leaf_len):
tree, leaf, _ = self._insert_and_lock(
window=window,
page_size=page_size,
leaf_len=leaf_len,
flag_on=True,
)
self.assertEqual(len(leaf.value), leaf_len)
self.assertEqual(tree.swa_protected_size_, leaf_len)
def test_match_prefix_returns_full_chain_after_split(self):
tree, allocator, _ = _build_swa_tree(
is_eagle=False,
kv_size=128,
kv_size_swa=64,
sliding_window_size=4,
page_size=1,
)
token_ids = list(range(12))
with envs.SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.override(True):
inserted_leaf = _insert_chain(tree, allocator, token_ids)
self.assertEqual(len(inserted_leaf.value), 4)
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", token_ids)))
)
self.assertEqual(match.device_indices.shape[0], 12)
self.assertIs(match.last_device_node, inserted_leaf)
def test_dec_lock_ref_after_split_balances_to_zero(self):
tree, leaf, result = self._insert_and_lock(
window=4, page_size=1, leaf_len=12, flag_on=True
)
self.assertEqual(tree.swa_protected_size_, 4)
self.assertEqual(tree.full_protected_size_, 12)
tree.dec_lock_ref(
leaf,
params=DecLockRefParams(swa_uuid_for_lock=result.swa_uuid_for_lock),
)
self.assertEqual(tree.swa_protected_size_, 0)
self.assertEqual(tree.full_protected_size_, 0)
tree.sanity_check()
class _SinglePoolAllocator(BaseTokenToKVPoolAllocator): class _SinglePoolAllocator(BaseTokenToKVPoolAllocator):
"""Minimal single-pool allocator: no SWA peer, so the whole range dies """Minimal single-pool allocator: no SWA peer, so the whole range dies
@@ -995,7 +372,7 @@ class TestFreeFullPartition(CustomTestCase):
"""`free_full` releases only the full side of a hybrid SWA allocator.""" """`free_full` releases only the full side of a hybrid SWA allocator."""
def setUp(self): def setUp(self):
_, self.allocator, _ = _build_swa_tree(is_eagle=False) self.allocator, _ = _build_swa_tree(is_eagle=False)
self.full_baseline = self.allocator.full_available_size() self.full_baseline = self.allocator.full_available_size()
self.swa_baseline = self.allocator.swa_available_size() self.swa_baseline = self.allocator.swa_available_size()
@@ -1043,7 +420,7 @@ class TestFreeKvRow(CustomTestCase):
whole, the SWA side only from the floor up.""" whole, the SWA side only from the floor up."""
def setUp(self): def setUp(self):
_, self.allocator, _ = _build_swa_tree(is_eagle=False) self.allocator, _ = _build_swa_tree(is_eagle=False)
self.full_baseline = self.allocator.full_available_size() self.full_baseline = self.allocator.full_available_size()
self.swa_baseline = self.allocator.swa_available_size() self.swa_baseline = self.allocator.swa_available_size()
@@ -1080,7 +457,7 @@ class TestFreeKvRow(CustomTestCase):
self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline)) self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline))
def test_below_floor_pieces_go_back_through_the_full_side(self): def test_below_floor_pieces_go_back_through_the_full_side(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) allocator, _ = _build_swa_tree(is_eagle=False, page_size=4)
indices = _swa_alloc(allocator, 8) indices = _swa_alloc(allocator, 8)
allocator.free_swa(indices) allocator.free_swa(indices)
after_alloc = allocator.full_available_size() after_alloc = allocator.full_available_size()
@@ -1100,7 +477,7 @@ class TestFreeKvRow(CustomTestCase):
self.assertEqual(allocator.full_available_size(), after_alloc + 8) self.assertEqual(allocator.full_available_size(), after_alloc + 8)
def test_grouped_full_side_frees_defer_and_skip_the_unique_path(self): def test_grouped_full_side_frees_defer_and_skip_the_unique_path(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) allocator, _ = _build_swa_tree(is_eagle=False, page_size=4)
indices = _swa_alloc(allocator, 12) indices = _swa_alloc(allocator, 12)
allocator.free_swa(indices[:8]) allocator.free_swa(indices[:8])
after_alloc = allocator.full_available_size() after_alloc = allocator.full_available_size()
@@ -1163,7 +540,7 @@ class TestSWAPeerMappedContract(CustomTestCase):
return bool(assert_async.call_args.args[0]) return bool(assert_async.call_args.args[0])
def test_segment_free_flags_a_page_whose_peer_is_already_gone(self): def test_segment_free_flags_a_page_whose_peer_is_already_gone(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) allocator, _ = _build_swa_tree(is_eagle=False, page_size=4)
live = _swa_alloc(allocator, 8) live = _swa_alloc(allocator, 8)
stale = _swa_alloc(allocator, 8) stale = _swa_alloc(allocator, 8)
allocator.clear_full_to_swa_mapping(stale) allocator.clear_full_to_swa_mapping(stale)
@@ -1176,7 +553,7 @@ class TestSWAPeerMappedContract(CustomTestCase):
"""page_size > 1: page reps by stride replace the page expansion's """page_size > 1: page reps by stride replace the page expansion's
filter and the inner allocator's torch.unique, in and out of a group.""" filter and the inner allocator's torch.unique, in and out of a group."""
ps = 4 ps = 4
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=ps) allocator, _ = _build_swa_tree(is_eagle=False, page_size=ps)
def grouped(indices): def grouped(indices):
allocator.free_group_begin() allocator.free_group_begin()
@@ -1204,7 +581,7 @@ class TestSWAPeerMappedContract(CustomTestCase):
self.assertIsNone(_sync_error(lambda: grouped(second[: 2 * ps - 1]))) self.assertIsNone(_sync_error(lambda: grouped(second[: 2 * ps - 1])))
def test_free_swa_flags_a_slot_whose_peer_is_already_gone(self): def test_free_swa_flags_a_slot_whose_peer_is_already_gone(self):
_, allocator, _ = _build_swa_tree(is_eagle=False) allocator, _ = _build_swa_tree(is_eagle=False)
live = _swa_alloc(allocator, 4) live = _swa_alloc(allocator, 4)
stale = _swa_alloc(allocator, 4) stale = _swa_alloc(allocator, 4)
# Whoever released the peer left the mapping reading as the padding slot. # Whoever released the peer left the mapping reading as the padding slot.
@@ -1217,7 +594,7 @@ class TestSWAPeerMappedContract(CustomTestCase):
def test_free_swa_does_not_synchronize(self): def test_free_swa_does_not_synchronize(self):
"""The filter's output shape was data-dependent, so it read a count back """The filter's output shape was data-dependent, so it read a count back
to the host; the gather that replaced it has a fixed shape.""" to the host; the gather that replaced it has a fixed shape."""
_, allocator, _ = _build_swa_tree(is_eagle=False) allocator, _ = _build_swa_tree(is_eagle=False)
mapping = allocator.full_to_swa_index_mapping mapping = allocator.full_to_swa_index_mapping
# Warm up outside the window: a first-time cudaMalloc can synchronize on # Warm up outside the window: a first-time cudaMalloc can synchronize on
@@ -1241,7 +618,7 @@ class TestSWAReqRingFree(CustomTestCase):
def _allocated_ring(self): def _allocated_ring(self):
ps = self.PS ps = self.PS
_, allocator, req_pool = _build_swa_tree( allocator, req_pool = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=ps, page_size=ps,
req_size=2, req_size=2,
@@ -1361,7 +738,7 @@ class TestSWAPageRepsFree(CustomTestCase):
PS = 4 PS = 4
def _allocator(self): def _allocator(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=self.PS) allocator, _ = _build_swa_tree(is_eagle=False, page_size=self.PS)
return allocator return allocator
def _sizes(self, allocator): def _sizes(self, allocator):
@@ -1371,7 +748,7 @@ class TestSWAPageRepsFree(CustomTestCase):
def test_free_swa_segment_npu_uses_reference_path(self): def test_free_swa_segment_npu_uses_reference_path(self):
for page_size in (1, 4): for page_size in (1, 4):
with self.subTest(page_size=page_size): with self.subTest(page_size=page_size):
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=page_size, page_size=page_size,
kv_size=8 * page_size, kv_size=8 * page_size,
@@ -1423,7 +800,7 @@ class TestSWAPageRepsFree(CustomTestCase):
with self.subTest(name=name): with self.subTest(name=name):
num_allocated_tokens = max(2 * page_size, num_tokens) num_allocated_tokens = max(2 * page_size, num_tokens)
kv_size = max(8 * page_size, num_allocated_tokens) kv_size = max(8 * page_size, num_allocated_tokens)
_, allocator, _ = _build_swa_tree( allocator, _ = _build_swa_tree(
is_eagle=False, is_eagle=False,
page_size=page_size, page_size=page_size,
kv_size=kv_size, kv_size=kv_size,
@@ -1467,90 +844,6 @@ class TestSWAPageRepsFree(CustomTestCase):
self.assertTrue(torch.all(mapping[indices[:touched]] == 0)) self.assertTrue(torch.all(mapping[indices[:touched]] == 0))
self.assertTrue(torch.all(mapping[indices[touched:]] > 0)) self.assertTrue(torch.all(mapping[indices[touched:]] > 0))
def test_node_frees_take_the_page_path_through_the_tree(self):
"""Tree values are page-aligned copies of a kv row, so SWA eviction and
the full eviction of its tombstones both free by page reps."""
ps = self.PS
tree, allocator, _ = _build_swa_tree(
is_eagle=False, page_size=ps, sliding_window_size=ps
)
full_before, swa_before = self._sizes(allocator)
_insert(tree, allocator, list(range(1, 3 * ps + 1)))
# Either inner `free` is the torch.unique path a caller falls back to
# when it hands no start position.
with (
patch.object(
allocator.full_attn_allocator,
"free",
side_effect=AssertionError("full side took the unique path"),
),
patch.object(
allocator.swa_attn_allocator,
"free",
side_effect=AssertionError("swa side took the unique path"),
),
):
tree.evict(EvictParams(num_tokens=0, swa_num_tokens=ps))
tree.evict(EvictParams(num_tokens=3 * ps, swa_num_tokens=0))
self.assertEqual(self._sizes(allocator), (full_before, swa_before))
class TestCacheUnfinishedReqEvictedPrefix(CustomTestCase):
"""An unfinished request whose SWA prefix is already gone must insert that
prefix as a tombstone, not as live SWA KV."""
def test_evicted_prefix_inserts_as_tombstone(self):
page_size, window, num_tokens, evicted = 4, 4, 16, 8
tree, allocator, req_to_token_pool = _build_swa_tree(
is_eagle=False, page_size=page_size, sliding_window_size=window
)
kv_indices = _swa_alloc(allocator, num_tokens)
req_to_token_pool.write((0, slice(0, num_tokens)), kv_indices)
# Drop the prefix's SWA peers, as window eviction would.
allocator.free_swa(kv_indices[:evicted])
swa_before = allocator.swa_available_size()
token_ids = array("q", range(1, num_tokens + 1))
req = _DummyReq()
req.kv.req_pool_idx = 0
req.origin_input_ids = token_ids
req.output_ids = array("q")
req.get_fill_ids = lambda: token_ids
req.extra_key = None
req.cache_salt = None
req.kv.cache_protected_len = 0
req.last_node = tree.root_node
req.lock_receipt = DecLockRefParams()
req.prefix_indices = torch.empty(0, dtype=torch.int64, device=tree.device)
req.kv.swa_evicted_seqlen = evicted
tree.cache_unfinished_req(req)
# The insert itself frees nothing.
self.assertEqual(allocator.swa_available_size(), swa_before)
# The live leaf holds a full window, so the whole key stays matchable.
self.assertEqual(req.kv.cache_protected_len, num_tokens)
# [0, evicted) is a tombstone; only [evicted, num_tokens) counts as SWA.
(first,) = tree.root_node.children.values()
self.assertTrue(first.swa_tombstone)
self.assertEqual(len(first.value), evicted)
self.assertEqual(
tree.swa_evictable_size_ + tree.swa_protected_size_,
num_tokens - evicted,
)
# Finishing drops the locks, which sanity_check needs; the accounting
# must survive the re-walk.
tree.cache_finished_req(req, owned_kv_len=num_tokens)
self.assertEqual(allocator.swa_available_size(), swa_before)
self.assertEqual(
tree.swa_evictable_size_ + tree.swa_protected_size_,
num_tokens - evicted,
)
tree.sanity_check()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -1,7 +1,7 @@
"""Large-scale benchmark + fuzz correctness tests for UnifiedRadixCache. """Large-scale benchmark + fuzz correctness tests for UnifiedRadixCache.
Usage (standalone): Usage (standalone):
bench: python3 test/registered/unit/mem_cache/test_unified_radix_cache_bench.py --bench --num-seqs 5000 --verify --components mamba legacy-mamba swa legacy-swa bench: python3 test/registered/unit/mem_cache/test_unified_radix_cache_bench.py --bench --num-seqs 5000 --verify --components mamba swa
CI Test: python -m pytest test/registered/unit/mem_cache/test_unified_radix_cache_bench.py -v -s CI Test: python -m pytest test/registered/unit/mem_cache/test_unified_radix_cache_bench.py -v -s
""" """
@@ -30,10 +30,8 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams, MatchPrefixParams,
) )
from sglang.srt.mem_cache.cache_init_params import CacheInitParams 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.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey 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.base import ComponentType from sglang.srt.mem_cache.unified_cache.components.base import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache 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.server_args import ServerArgs, set_global_server_args_for_scheduler
@@ -131,7 +129,6 @@ def create_bench_cache(
max_context_len, max_context_len,
components, components,
page_size=1, page_size=1,
tree_cls=None,
sliding_window_size=_SWA_WINDOW_SIZE, sliding_window_size=_SWA_WINDOW_SIZE,
): ):
"""Create cache. Returns (tree, allocator, req_to_token_pool, make_req).""" """Create cache. Returns (tree, allocator, req_to_token_pool, make_req)."""
@@ -225,21 +222,19 @@ def create_bench_cache(
) )
# --- tree --- # --- tree ---
if tree_cls is None:
tree_cls = UnifiedRadixCache
backend_override = ( backend_override = (
envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(_TREE_CORE_TEST_BACKEND) envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(_TREE_CORE_TEST_BACKEND)
if _TREE_CORE_TEST_BACKEND is not None and tree_cls is UnifiedRadixCache if _TREE_CORE_TEST_BACKEND is not None
else nullcontext() else nullcontext()
) )
with backend_override: with backend_override:
tree = tree_cls( tree = UnifiedRadixCache(
params=CacheInitParams( params=CacheInitParams(
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator, token_to_kv_pool_allocator=allocator,
page_size=page_size, page_size=page_size,
disable=False, disable=False,
tree_components=components if tree_cls is UnifiedRadixCache else None, tree_components=components,
sliding_window_size=sliding_window_size if has_swa else None, sliding_window_size=sliding_window_size if has_swa else None,
) )
) )
@@ -279,7 +274,7 @@ class _Env:
avg_tokens: int avg_tokens: int
def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None, page_size=1): def _make_env(num_seqs, chunk_len, kv_size, components, page_size=1):
"""Create sequences + cache, return shared _Env.""" """Create sequences + cache, return shared _Env."""
if components is None: if components is None:
components = _DEFAULT_COMPONENTS components = _DEFAULT_COMPONENTS
@@ -293,7 +288,6 @@ def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None, page_size
max_context_len=max_seq_len + 10, max_context_len=max_seq_len + 10,
components=components, components=components,
page_size=page_size, page_size=page_size,
tree_cls=tree_cls,
) )
return _Env( return _Env(
tree, tree,
@@ -467,11 +461,10 @@ def bench_insert(
kv_size=500_000, kv_size=500_000,
components=None, components=None,
verify=False, verify=False,
tree_cls=None,
page_size=1, page_size=1,
): ):
"""Insert throughput (alloc + evict-fallback + insert).""" """Insert throughput (alloc + evict-fallback + insert)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size) env = _make_env(num_seqs, chunk_len, kv_size, components, page_size)
warmup = min(20, num_seqs // 10) warmup = min(20, num_seqs // 10)
return bench_api( return bench_api(
@@ -491,11 +484,10 @@ def bench_match_prefix(
kv_size=500_000, kv_size=500_000,
components=None, components=None,
verify=False, verify=False,
tree_cls=None,
page_size=1, page_size=1,
): ):
"""Prefix matching throughput (hit / partial / miss mix).""" """Prefix matching throughput (hit / partial / miss mix)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size) env = _make_env(num_seqs, chunk_len, kv_size, components, page_size)
_populate(env, num_seqs // 2) _populate(env, num_seqs // 2)
rng = random.Random(123) rng = random.Random(123)
@@ -535,11 +527,10 @@ def bench_evict(
kv_size=500_000, kv_size=500_000,
components=None, components=None,
verify=False, verify=False,
tree_cls=None,
page_size=1, page_size=1,
): ):
"""Eviction throughput — fill pool then repeatedly evict batches.""" """Eviction throughput — fill pool then repeatedly evict batches."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size) env = _make_env(num_seqs, chunk_len, kv_size, components, page_size)
inserted = _fill_no_evict(env) inserted = _fill_no_evict(env)
evict_batch = max(100, kv_size // 200) evict_batch = max(100, kv_size // 200)
@@ -564,11 +555,10 @@ def bench_lock_unlock(
kv_size=500_000, kv_size=500_000,
components=None, components=None,
verify=False, verify=False,
tree_cls=None,
page_size=1, page_size=1,
): ):
"""Lock/unlock throughput — match nodes then cycle lock/unlock.""" """Lock/unlock throughput — match nodes then cycle lock/unlock."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls, page_size) env = _make_env(num_seqs, chunk_len, kv_size, components, page_size)
_populate(env, num_seqs // 2) _populate(env, num_seqs // 2)
nodes = [] nodes = []
@@ -608,14 +598,13 @@ def bench_cache_finished(
kv_size=500_000, kv_size=500_000,
components=None, components=None,
verify=False, verify=False,
tree_cls=None,
page_size=1, page_size=1,
): ):
"""cache_finished_req throughput — full request lifecycle. """cache_finished_req throughput — full request lifecycle.
Simulates: match_prefix → inc_lock_ref → alloc → fill req_to_token → cache_finished_req. 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, page_size) env = _make_env(num_seqs, chunk_len, kv_size, components, page_size)
# Pre-build Req objects with token IDs filled into req_to_token # Pre-build Req objects with token IDs filled into req_to_token
req_items: list = [] req_items: list = []
@@ -691,7 +680,6 @@ def run_all_benchmarks(
components=None, components=None,
verify=False, verify=False,
benchmarks=None, benchmarks=None,
tree_cls=None,
page_size=1, page_size=1,
): ):
if components is None: if components is None:
@@ -700,12 +688,12 @@ def run_all_benchmarks(
benchmarks = list(ALL_BENCHMARKS.keys()) benchmarks = list(ALL_BENCHMARKS.keys())
server_args = ServerArgs(model_path="dummy", page_size=page_size) server_args = ServerArgs(model_path="dummy", page_size=page_size)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise # The mamba component reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model. # loads the HF config for self.model_path — impossible for the dummy model.
server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size) server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size)
set_global_server_args_for_scheduler(server_args) set_global_server_args_for_scheduler(server_args)
impl_name = (tree_cls or UnifiedRadixCache).__name__ impl_name = UnifiedRadixCache.__name__
results = [] results = []
for name in benchmarks: for name in benchmarks:
if name not in ALL_BENCHMARKS: if name not in ALL_BENCHMARKS:
@@ -718,7 +706,6 @@ def run_all_benchmarks(
kv_size=kv_size, kv_size=kv_size,
components=components, components=components,
verify=verify, verify=verify,
tree_cls=tree_cls,
page_size=page_size, page_size=page_size,
) )
) )
@@ -823,12 +810,10 @@ del _cfg, _name
# CLI # CLI
# =================================================================== # ===================================================================
_TREE_CONFIGS = { _TREE_CONFIGS = {
"full": ((ComponentType.FULL,), None), "full": (ComponentType.FULL,),
"mamba": ((ComponentType.FULL, ComponentType.MAMBA), None), "mamba": (ComponentType.FULL, ComponentType.MAMBA),
"swa": ((ComponentType.FULL, ComponentType.SWA), None), "swa": (ComponentType.FULL, ComponentType.SWA),
"all": ((ComponentType.FULL, ComponentType.SWA, ComponentType.MAMBA), None), "all": (ComponentType.FULL, ComponentType.SWA, ComponentType.MAMBA),
"legacy-mamba": ((ComponentType.FULL, ComponentType.MAMBA), MambaRadixCache),
"legacy-swa": ((ComponentType.FULL, ComponentType.SWA), SWARadixCache),
} }
@@ -841,7 +826,7 @@ def _run_bench_cli():
"--components", "--components",
nargs="+", nargs="+",
choices=list(_TREE_CONFIGS.keys()), choices=list(_TREE_CONFIGS.keys()),
default=["mamba", "legacy-mamba"], default=["mamba"],
help="Component configs to benchmark", help="Component configs to benchmark",
) )
parser.add_argument("--page-size", type=int, default=1) parser.add_argument("--page-size", type=int, default=1)
@@ -857,7 +842,7 @@ def _run_bench_cli():
args, _ = parser.parse_known_args() args, _ = parser.parse_known_args()
for comp_name in args.components: for comp_name in args.components:
components, tree_cls = _TREE_CONFIGS[comp_name] components = _TREE_CONFIGS[comp_name]
run_all_benchmarks( run_all_benchmarks(
num_seqs=args.num_seqs, num_seqs=args.num_seqs,
chunk_len=args.chunk_len, chunk_len=args.chunk_len,
@@ -865,7 +850,6 @@ def _run_bench_cli():
components=components, components=components,
verify=args.verify, verify=args.verify,
benchmarks=args.benchmarks, benchmarks=args.benchmarks,
tree_cls=tree_cls,
page_size=args.page_size, page_size=args.page_size,
) )
@@ -471,7 +471,7 @@ def build_fixture(
page_size=cfg.page_size, page_size=cfg.page_size,
enable_int8_mamba_checkpoint=cfg.enable_int8_mamba_checkpoint, enable_int8_mamba_checkpoint=cfg.enable_int8_mamba_checkpoint,
) )
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise # The mamba component reads mamba_cache_chunk_size, whose property otherwise
# loads the HF config for self.model_path — impossible for the dummy model. # loads the HF config for self.model_path — impossible for the dummy model.
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE. # Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
server_args._mamba_cache_chunk_size = ( server_args._mamba_cache_chunk_size = (