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

Co-authored-by: ispobock <ispobaoke@gmail.com>
Co-authored-by: pansicheng <sicheng.pan.chn@gmail.com>
Co-authored-by: yizhang2077 <1109276519@qq.com>
Co-authored-by: xiezhq-hermann <xiezhq@stanford.edu>
This commit is contained in:
Zhangheng
2026-04-13 10:28:22 +08:00
committed by GitHub
co-authored by ispobock pansicheng yizhang2077 xiezhq-hermann
parent 5593539942
commit bc59cc0f96
15 changed files with 4707 additions and 1 deletions
+3
View File
@@ -472,6 +472,9 @@ class Envs:
SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16") SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16")
SGLANG_MAMBA_SSM_DTYPE = EnvStr(None) SGLANG_MAMBA_SSM_DTYPE = EnvStr(None)
# Unified Radix Tree
SGLANG_ENABLE_UNIFIED_RADIX_TREE = EnvBool(False)
# Breakable CUDA Graph # Breakable CUDA Graph
SGLANG_USE_BREAKABLE_CUDA_GRAPH = EnvBool(False) SGLANG_USE_BREAKABLE_CUDA_GRAPH = EnvBool(False)
+15
View File
@@ -835,6 +835,21 @@ class Scheduler(
self.tp_worker.register_hicache_layer_transfer_counter( self.tp_worker.register_hicache_layer_transfer_counter(
self.tree_cache.cache_controller.layer_done_counter self.tree_cache.cache_controller.layer_done_counter
) )
elif envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get():
from sglang.srt.mem_cache.unified_cache_components import (
ComponentType,
)
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
)
tree_components = [ComponentType.FULL]
if self.is_hybrid_swa or self.is_hybrid_ssm:
tree_components.append(
ComponentType.SWA if self.is_hybrid_swa else ComponentType.MAMBA
)
params.tree_components = tuple(tree_components)
self.tree_cache = UnifiedRadixCache(params)
elif self.is_hybrid_swa: elif self.is_hybrid_swa:
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
@@ -47,7 +47,7 @@ class MatchPrefixParams:
class InsertParams: class InsertParams:
"""Unified parameters for insert across different cache types""" """Unified parameters for insert across different cache types"""
key: RadixKey key: Optional[RadixKey] = None
value: Optional[torch.Tensor] = None value: Optional[torch.Tensor] = None
# Mamba specific # Mamba specific
@@ -8,6 +8,7 @@ import torch
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.unified_cache_components import ComponentType
@dataclasses.dataclass @dataclasses.dataclass
@@ -39,3 +40,5 @@ class CacheInitParams:
# Time-to-live for cache entries in seconds. If None, TTL is disabled. # Time-to-live for cache entries in seconds. If None, TTL is disabled.
cache_ttl_seconds: Optional[float] = None cache_ttl_seconds: Optional[float] = None
tree_components: Optional[tuple[ComponentType, ...]] = None
@@ -0,0 +1,329 @@
# Unified Radix Cache
A component-based, pluggable prefix cache framework for SGLang that unifies Full-attention, Sliding-Window-Attention (SWA), and Mamba/SSM caching into a single radix tree.
## Design Goals
1. **Unified tree structure** — One radix tree manages all KV cache types instead of separate specialized implementations (`SWARadixCache`, `MambaRadixCache`, etc.).
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 LRU list, lock reference counting, evictable/protected size tracking, and eviction driver.
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.
5. **Zero special-casing in the main tree** — The tree operates purely on keys (logical). All physical resource management (allocation, freeing, copy-on-write) is handled by components through hooks.
## Architecture
```
┌───────────────────────────────────────────────┐
│ UnifiedRadixCache │
│ (unified_radix_cache.py) │
│ │
│ root_node ──► UnifiedTreeNode (radix tree) │
│ components ► {name → TreeComponent} │
│ lru_lists ─► {name → UnifiedLRUList} │
└──────────┬───────────┬───────────┬────────────┘
│ │ │
▼ ▼ ▼
┌────────────┐ ┌──────────┐ ┌─────────────┐
│ Full │ │ SWA │ │ Mamba │
│ Component │ │Component │ │ Component │
└─────┬──────┘ └────┬─────┘ └──────┬──────┘
│ │ │
└─────────────┼──────────────┘
▼
┌──────────────┐
│TreeComponent │
│ (ABC) │
└──────────────┘
```
### Key Data Structures
**`UnifiedTreeNode`** — Each node stores per-component data independently:
```python
node.component_data = {
"full": ComponentData(value=Tensor|None, lock_ref=int, metadata={}),
"swa": ComponentData(value=Tensor|None, lock_ref=int, metadata={}),
"mamba": ComponentData(value=Tensor|None, lock_ref=int, metadata={}),
}
```
**`UnifiedLRUList`** — One doubly-linked list per component, threaded through the same tree nodes via `lru_prev[name]`/`lru_next[name]`. Supports O(1) insert/remove/promote and O(L) scan for eviction (L = locked nodes skipped).
**`ComponentData`** — Per-component data stored on each node:
- `value: Tensor | None` — Device indices into the component's memory pool (`TokenToKVPool` for Full, `SWAKVPool` for SWA, `MambaPool` for Mamba). `None` means tombstone (data evicted but node structure retained).
- `lock_ref: int` — Reference count of active requests using this node's component data. `lock_ref > 0` protects the node from eviction.
- `metadata: dict` — Component-specific state (e.g., SWA stores `component_uuid` for window-lock boundary tracking).
---
## File Layout
| File | Contents |
|------|----------|
| `../unified_radix_cache.py` | `UnifiedRadixCache`, `UnifiedTreeNode`, `UnifiedLRUList`, factory `create_unified_radix_cache` |
| `tree_component.py` | `TreeComponent` ABC, `ComponentType`, `ComponentData`, `get_and_increase_time_counter`, `next_component_uuid` |
| `full_component.py` | `FullComponent` — standard full-attention KV cache component |
| `swa_component.py` | `SWAComponent` — sliding-window attention component with tombstone/window tracking |
| `mamba_component.py` | `MambaComponent` — Mamba/SSM state component with copy-on-write |
| `hybrid_cache_controller.py` | `HybridCacheController` — HiCache 3-tier storage controller (L1 GPU → L2 CPU → L3 Disk) |
| `__init__.py` | Re-exports: `ComponentName`, `ComponentData`, `TreeComponent`, `FullComponent`, `SWAComponent`, `MambaComponent` |
---
## Public API Reference
All public APIs are on `UnifiedRadixCache`, which implements `BasePrefixCache`.
**Notation**: K = key length (tokens), D = matched path depth in tree (D ≤ K/P), P = page_size, C = number of components (≤ 3, treated as constant).
All tree traversal operations have two cost components: **O(K)** for data operations (key comparison, tensor clone/concat) + **O(D·C)** for component overhead (C hooks per node). Since D ≤ K/P and C is constant, overall **O(K)**.
### `match_prefix(params: MatchPrefixParams) → MatchResult`
Find the longest cached prefix for a token sequence.
| Aspect | Detail |
|--------|--------|
| **Purpose** | Walk the radix tree to find the longest prefix where **all** component validators pass |
| **Inputs** | `params.key: RadixKey` — token IDs + optional extra key for namespace isolation |
| **Output** | `MatchResult(device_indices, last_device_node, last_host_node, mamba_branching_seqlen, ...)` |
| **Mutation** | Updates `last_access_time` on matched path; promotes matched nodes to MRU in all component LRU lists; may trigger `_split_node` if match ends mid-node |
| **Complexity** | **O(K + D·C)** |
**Algorithm detail:**
1. Calls `create_match_validator()` once per component — returns a stateful closure (e.g., SWA tracks accumulated window length)
2. Walks tree edges via `key_match_fn`; at each node, calls all validator closures — the match boundary is only advanced when **all** validators return `True`
3. If match ends mid-node, calls `_split_node` → triggers `redistribute_on_node_split()` per component
4. Post-match (`_match_post_processor`):
- Promotes matched path to MRU in each component's LRU via `node_has_component_data()` as filter
- Updates `last_access_time` with decreasing timestamps up the path (parent < child)
- Concatenates matched device indices via `torch.cat` (concat length ≤ K, subsumed by O(K))
- Calls `finalize_match_result()` per component (Mamba performs copy-on-write: allocates new pool slot, copies SSM state)
---
### `insert(params: InsertParams) → InsertResult`
Insert a key-value pair into the tree.
| Aspect | Detail |
|--------|--------|
| **Purpose** | Insert token sequence + KV indices, reusing existing prefix and freeing duplicate KV slots |
| **Inputs** | `params.key: RadixKey`, `params.value: Tensor` (KV pool indices), plus component-specific fields (`mamba_value`, `swa_evicted_seqlen`, `prev_prefix_len`) |
| **Output** | `InsertResult(prefix_len, mamba_exist)` — `prefix_len` = length of reused prefix |
| **Mutation** | Creates new leaf nodes; updates component data on overlapping nodes; frees duplicate KV indices; may split nodes; updates LRU lists and evictable sizes |
| **Complexity** | **O(K + D·C)** |
**Algorithm detail** (`_insert_helper`):
1. At each existing node, calls `_touch_node` → promotes to MRU via `node_has_component_data()`
2. If key diverges mid-node, calls `_split_node` → `redistribute_on_node_split()` per component
3. For each overlapping node, calls `update_component_on_insert_overlap()` per component — returns `consumed_from` index; the tree frees `value[dup_start:consumed_from]` as duplicate pool indices
- Full: returns `prefix_len` (no consumption, default behavior)
- SWA: checks if the overlapping node is a tombstone (SWA value = None) within the SWA window boundary (`swa_evicted_seqlen`):
- If entirely within window: **recovers tombstone** — frees old `full_value`, clones `value_slice`, translates to SWA indices, inserts into SWA LRU (returns `0` = all consumed)
- If partially within window: **splits node** at boundary, recovers SWA on the window portion (returns `start_idx`)
- If entirely outside window: returns `prefix_len` (no consumption)
- Mamba: returns `prefix_len` (no consumption, default behavior)
4. Before creating a new leaf, checks `should_skip_leaf_creation()` per component — any veto aborts leaf creation and frees remaining value
5. Creates leaf via `_add_new_node` (clones value tensor, inserts into Full LRU)
6. Calls `commit_insert_component_data()` per component on the final target node (SWA may trigger a secondary split for window boundary; Mamba sets mamba pool indices and inserts into Mamba LRU)
---
### `evict(params: EvictParams) → EvictResult`
Free cached tokens to reclaim memory.
| Aspect | Detail |
|--------|--------|
| **Purpose** | Each component drives eviction from its own LRU list until its target is met |
| **Inputs** | `params.num_tokens` (full), `params.swa_num_tokens` (SWA), `params.mamba_num` (Mamba) |
| **Output** | `EvictResult(num_tokens_evicted, swa_num_tokens_evicted, mamba_num_evicted)` |
| **Mutation** | Frees pool indices; removes nodes from LRU lists; deletes leaf nodes from tree; cascades to lower-priority components; walks up parent chain to delete tombstone ancestors |
| **Complexity** | **O(E·H + L)** — E = nodes evicted, H = tombstone chain height, L = locked nodes skipped in LRU scan. |
**Algorithm detail:**
1. Calls `drive_eviction()` for each component:
- Full: scans Full LRU from tail, only evicts **leaf** nodes (`get_leaf_lru_no_lock` — **O(L)**); calls `evict_component()` to free pool indices
- SWA: scans SWA LRU from tail; **internal** nodes are tombstoned (evict SWA data, keep node), **leaf** nodes are fully deleted; both trigger cascade
- Mamba: scans Mamba LRU from tail; **internal** nodes are tombstoned, **leaf** nodes are fully deleted; both trigger cascade
2. After each node eviction, calls `_cascade_evict`:
- Queries `eviction_priority()` per component; evicts all with priority ≤ trigger's
- Calls `evict_component()` + `node_has_component_data()` for cascaded components
- For leaf: removes from parent, then `_iteratively_delete_tombstone_leaf` walks up **O(H)** ancestors
**Cascade eviction rules:**
- **Leaf nodes**: all priorities = 0 → evicting any cascades to all (node deleted)
- **Internal nodes**: Full(2) > SWA(1) > Mamba(0)
- Evicting Mamba: no cascade
- Evicting SWA: cascades to Mamba
- Evicting Full: cascades to SWA + Mamba
---
### `inc_lock_ref(node: UnifiedTreeNode) → IncLockRefResult`
Lock a node to protect it (and its ancestors) from eviction.
| Aspect | Detail |
|--------|--------|
| **Purpose** | Called when a request begins using a cached prefix — prevents eviction of nodes it depends on |
| **Inputs** | `node` — the last matched node (deepest) |
| **Output** | `IncLockRefResult(swa_uuid_for_lock)` |
| **Mutation** | Increments `lock_ref` per component along the path; moves tokens from evictable to protected size counters |
| **Complexity** | **O(D)** — Full: node to root; SWA: up to window boundary O(min(D, W)); Mamba: O(1).|
**Algorithm detail:** Calls `acquire_component_lock()` for each component.
| Component | Strategy |
|-----------|----------|
| Full | **Path-lock**: walks from node to root, `lock_ref += 1` on every ancestor. On first lock (`lock_ref: 0→1`), moves tokens from `component_evictable_size_` to `component_protected_size_`. |
| SWA | **Window-lock**: walks upward, accumulating SWA value lengths until `sliding_window_size` is filled. Records a `component_uuid` at the boundary node for `dec_lock_ref` to know where to stop. |
| Mamba | **Single-node lock**: only `lock_ref += 1` on the node itself (mamba state is per-leaf, not per-path). |
---
### `dec_lock_ref(node, params?) → DecLockRefResult`
Unlock a previously locked node path.
| Aspect | Detail |
|--------|--------|
| **Purpose** | Called when a request finishes — releases eviction protection |
| **Inputs** | `node`, optional `params.swa_uuid_for_lock` for SWA boundary detection |
| **Output** | `DecLockRefResult()` |
| **Mutation** | Decrements `lock_ref` per component; moves tokens from protected back to evictable when `lock_ref` reaches 0 |
| **Complexity** | **O(D)** — symmetric to `inc_lock_ref` |
**Algorithm detail:** Calls `release_component_lock()` for each component. Full walks to root; SWA walks up until matching `component_uuid`; Mamba decrements single node.
---
### `cache_finished_req(req: Req, is_insert: bool = True)`
Cache a completed request's KV data into the tree.
| Aspect | Detail |
|--------|--------|
| **Purpose** | After a request finishes, insert its token/KV data into the tree for future reuse |
| **Inputs** | `req` — the finished request; `is_insert` — whether to insert (True) or just release locks (False) |
| **Output** | `None` |
| **Mutation** | Calls component hooks → `insert` → `dec_lock_ref` → component cleanup. Frees unaligned tail KV indices; frees non-inserted KV indices when `is_insert=False`. |
| **Complexity** | **O(K + D·C)** — insert O(K + D·C) + lock release O(D). Simplifies to **O(K)**. |
**Algorithm detail:**
1. `prepare_for_caching_req()` per component — sets component-specific insert params, returns effective cache length (SWA: sets `swa_evicted_seqlen`; Mamba: prepares `mamba_value` from ping-pong buffer, returns `mamba_last_track_seqlen` as truncation hint)
2. Truncates if `effective_cache_len < len(token_ids)`: frees excess pool indices
3. Converts token IDs (bigram if EAGLE), page-aligns keys, then calls `insert()`
4. Frees unaligned tail KV indices beyond page boundary
5. Calls `dec_lock_ref()` on the previous `req.last_node`
6. `cleanup_after_caching_req()` per component (Mamba: frees forked mamba_value based on `mamba_exist`, handles ping-pong buffer cleanup)
---
### `cache_unfinished_req(req: Req, chunked=False)`
Cache an in-progress request's partial KV data (chunked prefill).
| Aspect | Detail |
|--------|--------|
| **Purpose** | During chunked prefill, insert partial results so the next chunk can match the prefix |
| **Inputs** | `req` — the in-progress request |
| **Output** | `None` |
| **Mutation** | Inserts partial KV → re-matches prefix → updates `req.prefix_indices`, `req.cache_protected_len`, `req.last_node`; transfers lock from old node to new node |
| **Complexity** | **O(K + D·C)** — two tree traversals: insert O(K + D·C) + re-match O(K + D·C) + lock transfer O(D). Simplifies to **O(K)**. |
**Algorithm detail:**
1. `prepare_for_caching_req()` per component
2. `insert()` — first tree traversal
3. `match_prefix()` — **second** tree traversal to get updated indices
4. Writes new prefix indices into `req_to_token_pool`
5. `dec_lock_ref()` on old `req.last_node`
6. `inc_lock_ref()` on new matched node
7. Updates `req.prefix_indices`, `req.cache_protected_len`, `req.last_node`
8. `cleanup_after_caching_req()` per component
---
## TreeComponent Hook Reference
Each component implements these hooks. See `tree_component.py` for the ABC and docstrings.
### Match Phase
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `create_match_validator()` | Return a per-match stateful predicate that decides whether a node is a valid match boundary. Full: always True. SWA: tracks accumulated window length, True when contiguous window ≥ `sliding_window_size`. Mamba: True iff node has mamba data. | `_match_prefix_helper` | *abstract* |
| `finalize_match_result()` | Post-process the match result after prefix matching completes. Full/SWA: pass-through. Mamba: copy-on-write — allocates a new mamba pool slot, copies SSM state into the request pool, records `branching_seqlen`. | `_match_post_processor` | pass-through |
### Insert Phase
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `update_component_on_insert_overlap()` | Handle key overlap with an existing node during insert. Returns the index within `value_slice` from which this component consumed (took ownership of) pool slots. Full/Mamba: no consumption (`prefix_len`). SWA: may recover tombstoned nodes within the sliding window boundary. | `_insert_helper` | returns `prefix_len` |
| `should_skip_leaf_creation()` | Veto leaf creation when the entire new leaf would be a tombstone for this component. SWA: vetoes if `swa_evicted_seqlen ≥ total_prefix_len + key_len`. | `_insert_helper` | `False` |
| `commit_insert_component_data()` | Finalize component data on the target node after the insert walk completes. Full: no-op (handled by `_add_new_node`). SWA: checks window boundary, may split node — parent becomes tombstone, child gets SWA data. Mamba: sets mamba pool indices and inserts into Mamba LRU. | `_insert_helper` | no-op |
### Node Split
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `redistribute_on_node_split()` | Redistribute component data between new parent (prefix) and child (suffix) when a node is split. Full: copies `lock_ref` to parent. SWA: slices SWA value, copies `lock_ref` and `component_uuid`. Mamba: parent gets `None`/`lock_ref=0` (mamba stays on leaf). | `_split_node` | *abstract* |
### Eviction Phase
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `evict_component()` | Free this component's KV resources on a node being evicted. Internal nodes: free memory and tombstone (`value = None`). Leaf nodes: free memory, node will be deleted. Returns number of tokens freed. | `_evict_component_and_detach_lru` | *abstract* |
| `eviction_priority()` | Return cascade eviction priority (higher = evicted later). Leaf: all 0. Internal: Full(2) > SWA(1) > Mamba(0). When evicting, all components with ≤ priority on the same node are cascade-evicted. | `_cascade_evict` | `0` |
| `drive_eviction()` | Drive eviction from this component's LRU list until the target amount is freed. Full: leaf-only from Full LRU. SWA: both internal (tombstone) and leaf from SWA LRU. Mamba: both internal (tombstone) and leaf from Mamba LRU. | `evict` | *abstract* |
### Lock Phase
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `acquire_component_lock()` | Increment `lock_ref` to protect nodes from eviction; moves tokens from evictable to protected. Full: path-lock to root. SWA: window-lock with UUID boundary. Mamba: single-node lock. | `inc_lock_ref` | *abstract* |
| `release_component_lock()` | Decrement `lock_ref` to un-protect nodes; moves tokens from protected to evictable when `lock_ref` → 0. Full: path-unlock to root. SWA: walks up to UUID boundary. Mamba: single-node unlock. | `dec_lock_ref` | *abstract* |
### Caching Phase
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `prepare_for_caching_req()` | Prepare component-specific data before insert, fill fields in `InsertParams`, return effective cache length. Full: no-op. SWA: sets `swa_evicted_seqlen`. Mamba: prepares `mamba_value` from ping-pong buffer, returns `mamba_last_track_seqlen`. | `cache_finished/unfinished_req` | returns `None` |
| `cleanup_after_caching_req()` | Post-cache cleanup. Full/SWA: no-op. Mamba: frees forked `mamba_value` based on `mamba_exist`, handles ping-pong buffer `keep_idx`, resets `mamba_last_track_seqlen` on unfinished. | `cache_finished/unfinished_req` | no-op |
### Utility
| Hook | Purpose | Called By | Default |
|------|---------|-----------|----------|
| `node_has_component_data()` | Check if a node has this component's data. Used as filter for LRU operations and cascade checks. Full overrides to check `full_value` directly. | multiple | `value is not None` |
### Component Behavior Summary
| Behavior | FullComponent | SWAComponent | MambaComponent |
|----------|--------------|-------------|----------------|
| **Validator** | Always `True` | Tracks accumulated window; `True` when ≥ `sliding_window_size` | `True` iff node has mamba data |
| **Lock strategy** | Path-lock (root → node) | Window-lock (up to window boundary, UUID-tagged) | Single-node lock |
| **Internal eviction priority** | 2 (last) | 1 (middle) | 0 (first) |
| **Split behavior** | Copy `lock_ref` to parent | Slice SWA value + copy UUID | Parent gets `None` (mamba stays on leaf) |
| **Match finalize** | No-op | No-op | Copy-on-write: allocate new mamba slot, copy state |
| **Drive eviction** | Full LRU (leaf-only) → cascade all | SWA LRU → tombstone internal, cascade leaf | Mamba LRU → tombstone internal, cascade leaf |
---
## Factory Function
```python
def create_unified_radix_cache(
params: CacheInitParams,
component_names: Optional[tuple[ComponentName, ...]] = None,
) -> UnifiedRadixCache
```
Auto-detects component configuration from `params` if `component_names` is not specified:
- `SWATokenToKVPoolAllocator` → `(SWA,)` → `UnifiedSWARadixCache`
- `HybridReqToTokenPool` → `(MAMBA,)` → `UnifiedMambaRadixCache`
- Explicit tuple → `UnifiedRadixCache` with specified components
Enable via `--enable-unified-radix-tree` server flag.
@@ -0,0 +1,25 @@
from sglang.srt.mem_cache.unified_cache_components.full_component import FullComponent
from sglang.srt.mem_cache.unified_cache_components.mamba_component import MambaComponent
from sglang.srt.mem_cache.unified_cache_components.swa_component import SWAComponent
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
_NUM_COMPONENT_TYPES,
BASE_COMPONENT_TYPE,
ComponentData,
ComponentType,
TreeComponent,
get_and_increase_time_counter,
next_component_uuid,
)
__all__ = [
"BASE_COMPONENT_TYPE",
"ComponentData",
"ComponentType",
"FullComponent",
"MambaComponent",
"SWAComponent",
"TreeComponent",
"_NUM_COMPONENT_TYPES",
"next_component_uuid",
"get_and_increase_time_counter",
]
@@ -0,0 +1,102 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable, Optional
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
)
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
ComponentType,
TreeComponent,
)
if TYPE_CHECKING:
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedTreeNode,
)
class FullComponent(TreeComponent):
component_type = ComponentType.FULL
def __init__(self, cache, params):
super().__init__(cache, params)
allocator = cache.token_to_kv_pool_allocator
# When SWA is present, only free full-attention KV here;
# SWA KV will be freed by cascade via SWAComponent.evict_component.
if ComponentType.SWA in cache.tree_components:
self._free_full = allocator.full_attn_allocator.free
else:
self._free_full = allocator.free
def node_has_component_data(self, node: UnifiedTreeNode) -> bool:
# Override so _for_each_component_lru includes Full in LRU operations
return node.component_data[self.component_type].value is not None
def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]:
return lambda node: True
def redistribute_on_node_split(
self, new_parent: UnifiedTreeNode, child: UnifiedTreeNode
):
new_parent.component_data[self.component_type].lock_ref = child.component_data[
self.component_type
].lock_ref
def evict_component(self, node: UnifiedTreeNode, is_leaf: bool) -> int:
cd = node.component_data[self.component_type]
self._free_full(cd.value)
freed = len(cd.value)
self.cache.component_evictable_size_[self.component_type] -= freed
return freed
def eviction_priority(self, is_leaf: bool) -> int:
return 0 if is_leaf else 2
def drive_eviction(
self, params: EvictParams, tracker: dict[ComponentType, int]
) -> None:
request = params.num_tokens
lru = self.cache.lru_lists[self.component_type]
while tracker[self.component_type] < request:
x = lru.get_leaf_lru_no_lock()
if x is None:
break
self.cache._evict_component_and_detach_lru(
x, self, is_leaf=True, tracker=tracker
)
self.cache._cascade_evict(x, self, tracker)
def acquire_component_lock(
self, node: UnifiedTreeNode, result: IncLockRefResult
) -> IncLockRefResult:
ct = self.component_type
root = self.cache.root_node
cur = node
while cur != root:
cd = cur.component_data[ct]
if cd.lock_ref == 0:
key_len = len(cd.value)
self.cache.component_evictable_size_[ct] -= key_len
self.cache.component_protected_size_[ct] += key_len
cd.lock_ref += 1
cur = cur.parent
return result
def release_component_lock(
self, node: UnifiedTreeNode, params: Optional[DecLockRefParams]
) -> None:
ct = self.component_type
root = self.cache.root_node
cur = node
while cur != root:
cd = cur.component_data[ct]
assert cd.lock_ref > 0
if cd.lock_ref == 1:
key_len = len(cd.value)
self.cache.component_evictable_size_[ct] += key_len
self.cache.component_protected_size_[ct] -= key_len
cd.lock_ref -= 1
cur = cur.parent
@@ -0,0 +1,270 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable, Optional
import torch
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
InsertParams,
InsertResult,
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
ComponentType,
TreeComponent,
get_and_increase_time_counter,
)
from sglang.srt.server_args import get_global_server_args
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
UnifiedTreeNode,
)
class MambaComponent(TreeComponent):
component_type = ComponentType.MAMBA
def __init__(self, cache: UnifiedRadixCache, params: CacheInitParams):
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
assert isinstance(
cache.req_to_token_pool, HybridReqToTokenPool
), f"MambaComponent requires HybridReqToTokenPool, got {type(cache.req_to_token_pool)}"
if not params.enable_mamba_extra_buffer:
assert (
cache.page_size == 1
), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {cache.page_size}"
super().__init__(cache, params)
self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer
def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]:
ct = self.component_type
return lambda node: node.component_data[ct].value is not None
def finalize_match_result(
self,
result: MatchResult,
params: MatchPrefixParams,
value_chunks: list[torch.Tensor],
best_value_len: int,
) -> MatchResult:
cow_mamba = params.cow_mamba
req = params.req
last_node = result.last_device_node
if len(value_chunks) > best_value_len:
chunk_size = get_global_server_args().mamba_cache_chunk_size
aligned_seqlen = (
sum(len(v) for v in value_chunks) // chunk_size
) * chunk_size
branching_seqlen = aligned_seqlen if aligned_seqlen > 0 else None
else:
branching_seqlen = None
mamba_value = last_node.component_data[self.component_type].value
if cow_mamba and mamba_value is not None:
assert req is not None
if req.mamba_pool_idx is None:
dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1)
if dst_index is None:
self.cache.inc_lock_ref(last_node)
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
dst_index = self.cache.req_to_token_pool.mamba_pool.alloc(1)
self.cache.dec_lock_ref(last_node)
assert dst_index is not None, "Can not alloc mamba cache"
self.cache.req_to_token_pool.mamba_pool.copy_from(
mamba_value, dst_index
)
req.mamba_pool_idx = dst_index[0]
else:
dst_index = req.mamba_pool_idx.unsqueeze(0)
self.cache.req_to_token_pool.mamba_pool.copy_from(
mamba_value, dst_index
)
return result._replace(mamba_branching_seqlen=branching_seqlen)
def commit_insert_component_data(
self,
node: UnifiedTreeNode,
is_new_leaf: bool,
params: InsertParams,
result: InsertResult,
) -> None:
assert params.mamba_value is not None
if is_new_leaf:
node.component_data[self.component_type].value = params.mamba_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(
params.mamba_value
)
return
if node.component_data[self.component_type].value is None:
node.component_data[self.component_type].value = params.mamba_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(
params.mamba_value
)
node.last_access_time = get_and_increase_time_counter()
return
self.cache.lru_lists[self.component_type].reset_node_mru(node)
node.last_access_time = get_and_increase_time_counter()
result.mamba_exist = True
def redistribute_on_node_split(
self, new_parent: UnifiedTreeNode, child: UnifiedTreeNode
):
new_parent.component_data[self.component_type].value = None
new_parent.component_data[self.component_type].lock_ref = 0
def evict_component(self, node: UnifiedTreeNode, is_leaf: bool) -> int:
value = node.component_data[self.component_type].value
self.cache.req_to_token_pool.mamba_pool.free(value)
freed = len(value)
self.cache.component_evictable_size_[self.component_type] -= freed
if not is_leaf:
node.component_data[self.component_type].value = None
return freed
def drive_eviction(
self, params: EvictParams, tracker: dict[ComponentType, int]
) -> None:
request = params.mamba_num
lru = self.cache.lru_lists[self.component_type]
x = lru.get_lru_no_lock()
while (
tracker[self.component_type] < request and x is not None and lru.in_list(x)
):
assert x.component_data[self.component_type].value is not None
if len(x.children) > 0:
x_next = lru.get_prev_no_lock(x)
self.cache._evict_component_and_detach_lru(
x, self, is_leaf=False, tracker=tracker
)
self.cache._cascade_evict(x, self, tracker)
x = x_next
else:
self.cache._evict_component_and_detach_lru(
x, self, is_leaf=True, tracker=tracker
)
self.cache._cascade_evict(x, self, tracker)
x = lru.get_lru_no_lock()
def acquire_component_lock(
self, node: UnifiedTreeNode, result: IncLockRefResult
) -> IncLockRefResult:
ct = self.component_type
cd = node.component_data[ct]
value = cd.value
if value is not None:
if cd.lock_ref == 0:
vlen = len(value)
self.cache.component_evictable_size_[ct] -= vlen
self.cache.component_protected_size_[ct] += vlen
cd.lock_ref += 1
return result
def release_component_lock(
self, node: UnifiedTreeNode, params: Optional[DecLockRefParams]
) -> None:
ct = self.component_type
cd = node.component_data[ct]
value = cd.value
if value is not None:
assert cd.lock_ref > 0
if cd.lock_ref == 1:
vlen = len(value)
self.cache.component_evictable_size_[ct] += vlen
self.cache.component_protected_size_[ct] -= vlen
cd.lock_ref -= 1
def prepare_for_caching_req(
self,
req: Req,
insert_params: InsertParams,
token_ids_len: int,
is_finished: bool,
) -> Optional[int]:
cache_len = (
req.mamba_last_track_seqlen
if self.enable_mamba_extra_buffer
else token_ids_len
)
if is_finished:
if cache_len is None:
cache_len = 0
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
mamba_value = (
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
)
else:
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
insert_params.mamba_value = mamba_value
return cache_len
else:
if cache_len is None:
return 0
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
mamba_value = (
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
)
else:
mamba_value = self.cache.req_to_token_pool.get_mamba_indices(
req.req_pool_idx
).unsqueeze(-1)
mamba_value_forked = self.cache.req_to_token_pool.mamba_pool.fork_from(
mamba_value
)
if mamba_value_forked is None:
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
mamba_value_forked = self.cache.req_to_token_pool.mamba_pool.fork_from(
mamba_value
)
assert mamba_value_forked is not None, "Can not alloc mamba cache"
insert_params.mamba_value = mamba_value_forked
return cache_len
def cleanup_after_caching_req(
self,
req: Req,
is_finished: bool,
insert_result: Optional[InsertResult] = None,
insert_params: Optional[InsertParams] = None,
) -> None:
if is_finished:
mamba_exist = (
insert_result.mamba_exist if insert_result is not None else True
)
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx
)
else:
keep_idx = None
if mamba_exist:
keep_idx = None
free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist
if free_mamba_cache:
self.cache.req_to_token_pool.free_mamba_cache(
req, mamba_ping_pong_track_buffer_to_keep=keep_idx
)
else:
if insert_params.mamba_value is not None and (
insert_result is None or insert_result.mamba_exist
):
self.cache.req_to_token_pool.mamba_pool.free(insert_params.mamba_value)
req.mamba_last_track_seqlen = None
@@ -0,0 +1,297 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable, Optional
import torch
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
InsertParams,
InsertResult,
)
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
BASE_COMPONENT_TYPE,
ComponentType,
TreeComponent,
next_component_uuid,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
UnifiedTreeNode,
)
class SWAComponent(TreeComponent):
"""Sliding window attention component.
Each SWA node stores translated SWA pool indices as its component
value, independent of the full attention indices on the same tree node.
When SWA data is evicted from an internal node the node is tombstoned
— its SWA component value becomes None while the full attention
value stays intact.
"""
def __init__(self, cache: UnifiedRadixCache, params: CacheInitParams):
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
assert isinstance(
cache.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
), f"SWAComponent requires SWATokenToKVPoolAllocator, got {type(cache.token_to_kv_pool_allocator)}"
super().__init__(cache, params)
self.sliding_window_size = params.sliding_window_size
component_type = ComponentType.SWA
def _translate_full_to_swa(self, full_indices: torch.Tensor) -> torch.Tensor:
return self.cache.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
full_indices
)
def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]:
sliding_window_size = self.sliding_window_size
ct = self.component_type
state = {"len": float("inf")}
def validator(node: UnifiedTreeNode) -> bool:
if node.component_data[ct].value is None:
state["len"] = 0
return False
state["len"] += len(node.key)
return state["len"] >= sliding_window_size
return validator
def update_component_on_insert_overlap(
self,
node: UnifiedTreeNode,
prefix_len: int,
total_prefix_len: int,
value_slice: torch.Tensor,
params: InsertParams,
) -> int:
if params.prev_prefix_len >= total_prefix_len + prefix_len:
return prefix_len
is_tombstone = node.component_data[self.component_type].value is None
if not is_tombstone:
return prefix_len
swa_evicted_seqlen = params.swa_evicted_seqlen
assert (
node.component_data[self.component_type].lock_ref == 0
), f"tombstone {self.component_type} lock_ref should be 0, node {node.id}"
assert (
swa_evicted_seqlen % self.cache.page_size == 0
), f"{self.component_type}: swa_evicted_seqlen must be page-aligned, {swa_evicted_seqlen=}"
if swa_evicted_seqlen <= total_prefix_len:
# Branch 1: entire value_slice is within SWA window — recover
self.cache.token_to_kv_pool_allocator.free(
node.component_data[BASE_COMPONENT_TYPE].value
)
node.component_data[BASE_COMPONENT_TYPE].value = value_slice.clone()
swa_value = self._translate_full_to_swa(
node.component_data[BASE_COMPONENT_TYPE].value
)
node.component_data[self.component_type].value = swa_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(swa_value)
return 0
elif swa_evicted_seqlen < total_prefix_len + prefix_len:
# Branch 2: value_slice[start_idx:] is within SWA window — partial recover
start_idx = swa_evicted_seqlen - total_prefix_len
self.cache.token_to_kv_pool_allocator.free(
node.component_data[BASE_COMPONENT_TYPE].value[start_idx:]
)
self.cache._split_node(node.key, node, start_idx)
node.component_data[BASE_COMPONENT_TYPE].value = value_slice[
start_idx:
].clone()
swa_value = self._translate_full_to_swa(
node.component_data[BASE_COMPONENT_TYPE].value
)
node.component_data[self.component_type].value = swa_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(swa_value)
return start_idx
else:
# Branch 3: entire value_slice is outside SWA window — not consumed
return prefix_len
def should_skip_leaf_creation(
self, total_prefix_len: int, key_len: int, params: InsertParams
) -> bool:
return params.swa_evicted_seqlen >= total_prefix_len + key_len
def commit_insert_component_data(
self,
node: UnifiedTreeNode,
is_new_leaf: bool,
params: InsertParams,
result: InsertResult,
) -> None:
if not is_new_leaf:
return
node_start = result.prefix_len
split_pos = params.swa_evicted_seqlen - node_start
if split_pos <= 0:
swa_value = self._translate_full_to_swa(
node.component_data[BASE_COMPONENT_TYPE].value
)
node.component_data[self.component_type].value = swa_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(swa_value)
elif split_pos < len(node.key):
# Node straddles the SWA eviction boundary
# Split into parent (tombstone, no SWA) and child (with SWA)
# After _split_node, `node` becomes the child
self.cache._split_node(node.key, node, split_pos)
swa_value = self._translate_full_to_swa(
node.component_data[BASE_COMPONENT_TYPE].value
)
node.component_data[self.component_type].value = swa_value
self.cache.lru_lists[self.component_type].insert_mru(node)
self.cache.component_evictable_size_[self.component_type] += len(swa_value)
def redistribute_on_node_split(
self, new_parent: UnifiedTreeNode, child: UnifiedTreeNode
):
new_parent.component_data[self.component_type].lock_ref = child.component_data[
self.component_type
].lock_ref
child_swa_value = child.component_data[self.component_type].value
if child_swa_value is not None:
split_len = len(new_parent.key)
new_parent.component_data[self.component_type].value = child_swa_value[
:split_len
].clone()
child.component_data[self.component_type].value = child_swa_value[
split_len:
].clone()
else:
new_parent.component_data[self.component_type].value = None
# parent inherits the swa_uuid from child for swa lock ref
new_parent.component_data[self.component_type].metadata["uuid"] = (
child.component_data[self.component_type].metadata.get("uuid")
)
child.component_data[self.component_type].metadata.pop("uuid", None)
def evict_component(self, node: UnifiedTreeNode, is_leaf: bool) -> int:
swa_value = node.component_data[self.component_type].value
if swa_value is None:
return 0
# Direct swa_attn_allocator.free(swa_value) would double-free
# free_swa(full_value) has the mapping guard to avoid double-free
# TODO: decoupling full and swa free, need further discussion on mapping necessity
self.cache.token_to_kv_pool_allocator.free_swa(
node.component_data[BASE_COMPONENT_TYPE].value
)
freed = len(swa_value)
self.cache.component_evictable_size_[self.component_type] -= freed
if not is_leaf:
node.component_data[self.component_type].value = None
return freed
def eviction_priority(self, is_leaf: bool) -> int:
return 0 if is_leaf else 1
def drive_eviction(
self, params: EvictParams, tracker: dict[ComponentType, int]
) -> None:
request = params.swa_num_tokens
lru = self.cache.lru_lists[self.component_type]
x = lru.get_lru_no_lock()
while (
tracker[self.component_type] < request and x is not None and lru.in_list(x)
):
assert x.component_data[self.component_type].value is not None
if len(x.children) > 0:
x_next = lru.get_prev_no_lock(x)
self.cache._evict_component_and_detach_lru(
x, self, is_leaf=False, tracker=tracker
)
self.cache._cascade_evict(x, self, tracker)
x = x_next
else:
self.cache._evict_component_and_detach_lru(
x, self, is_leaf=True, tracker=tracker
)
self.cache._cascade_evict(x, self, tracker)
x = lru.get_lru_no_lock()
def acquire_component_lock(
self, node: UnifiedTreeNode, result: IncLockRefResult
) -> IncLockRefResult:
ct = self.component_type
root = self.cache.root_node
sliding_window_size = self.sliding_window_size
swa_lock_size = 0
swa_uuid_for_lock = None
cur = node
while cur != root and swa_lock_size < sliding_window_size:
assert (
cur.component_data[ct].value is not None
), f"acquire_component_lock({ct}) on tombstoned node {cur.id}"
comp = cur.component_data[ct]
if comp.lock_ref == 0:
key_len = len(cur.key)
self.cache.component_evictable_size_[ct] -= key_len
self.cache.component_protected_size_[ct] += key_len
comp.lock_ref += 1
swa_lock_size += len(cur.key)
if swa_lock_size >= sliding_window_size:
if comp.metadata.get("uuid") is None:
comp.metadata["uuid"] = next_component_uuid()
swa_uuid_for_lock = comp.metadata["uuid"]
cur = cur.parent
result.swa_uuid_for_lock = swa_uuid_for_lock
return result
def release_component_lock(
self, node: UnifiedTreeNode, params: Optional[DecLockRefParams]
) -> None:
ct = self.component_type
root = self.cache.root_node
swa_uuid_for_lock = params.swa_uuid_for_lock if params else None
dec_swa = True
cur = node
while cur != root and dec_swa:
assert (
cur.component_data[ct].value is not None
), f"release_component_lock({ct}) on tombstoned node {cur.id}"
comp = cur.component_data[ct]
assert (
comp.lock_ref > 0
), f"release_component_lock({ct}) on node with lock_ref=0, node {cur.id}"
if comp.lock_ref == 1:
key_len = len(cur.key)
self.cache.component_evictable_size_[ct] += key_len
self.cache.component_protected_size_[ct] -= key_len
comp.lock_ref -= 1
if swa_uuid_for_lock and comp.metadata.get("uuid") == swa_uuid_for_lock:
dec_swa = False
cur = cur.parent
def prepare_for_caching_req(
self,
req: Req,
insert_params: InsertParams,
token_ids_len: int,
is_finished: bool,
) -> Optional[int]:
if is_finished:
insert_params.swa_evicted_seqlen = req.swa_evicted_seqlen
return None
@@ -0,0 +1,292 @@
from __future__ import annotations
import dataclasses
from abc import ABC, abstractmethod
from enum import Enum
from typing import TYPE_CHECKING, Any, Callable, Optional
import torch
from numpy import float64
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
InsertParams,
InsertResult,
MatchPrefixParams,
MatchResult,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
UnifiedTreeNode,
)
class ComponentType(int, Enum):
"""Integer enum so that per-node list/tuple storage can be indexed directly."""
FULL = 0
SWA = 1
MAMBA = 2
def __str__(self) -> str: # keep human-readable logging
return self.name.lower()
@property
def is_full(self) -> bool:
return self == ComponentType.FULL
@property
def is_swa(self) -> bool:
return self == ComponentType.SWA
@property
def is_mamba(self) -> bool:
return self == ComponentType.MAMBA
BASE_COMPONENT_TYPE = ComponentType.FULL
_NUM_COMPONENT_TYPES = len(ComponentType)
_LAST_ACCESS_TIME_COUNTER_FLOAT = float64(1.0)
_COMPONENT_UUID_COUNTER = 1
@dataclasses.dataclass
class ComponentData:
value: Optional[torch.Tensor] = None
lock_ref: int = 0
metadata: dict[str, Any] = dataclasses.field(default_factory=dict)
def get_and_increase_time_counter() -> float64:
global _LAST_ACCESS_TIME_COUNTER_FLOAT
ret = _LAST_ACCESS_TIME_COUNTER_FLOAT
_LAST_ACCESS_TIME_COUNTER_FLOAT += 1.0
return ret
def next_component_uuid() -> int:
global _COMPONENT_UUID_COUNTER
_COMPONENT_UUID_COUNTER += 1
return _COMPONENT_UUID_COUNTER
class TreeComponent(ABC):
def __init__(self, cache: UnifiedRadixCache, params: CacheInitParams):
self.cache = cache
# Subclasses MUST set this as a class attribute (not @property)
component_type: ComponentType
def node_has_component_data(self, node: UnifiedTreeNode) -> bool:
return node.component_data[self.component_type].value is not None
def value_len(self, node: UnifiedTreeNode) -> int:
value = node.component_data[self.component_type].value
return len(value) if value is not None else 0
@abstractmethod
def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]:
"""Return a per-match stateful predicate that decides whether a node
is a valid match boundary for this component.
Called once per match_prefix; the returned closure may carry state.
- Full: always True (every node is valid).
- SWA: tracks accumulated length since last gap; returns True only
when the contiguous window reaches swa_sliding_window_size.
- Mamba: returns True iff the node has mamba component data."""
...
def finalize_match_result(
self,
result: MatchResult,
params: MatchPrefixParams,
value_chunks: list[torch.Tensor],
best_value_len: int,
) -> MatchResult:
"""Post-process the match result after prefix matching completes.
- Full & SWA: pass through unchanged.
- Mamba: performs copy-on-write — allocates a new mamba slot, copies
the matched node's mamba state into the request pool, and records
branching_seqlen in result."""
return result
def update_component_on_insert_overlap(
self,
node: UnifiedTreeNode,
prefix_len: int,
total_prefix_len: int,
value_slice: torch.Tensor,
params: InsertParams,
) -> int:
"""Called per-node when an insert's key overlaps an existing node.
Returns the index within value_slice from which this component
consumed (took ownership of) the underlying KV pool slots.
Returns prefix_len if nothing was consumed (default).
_insert_helper uses this to free only the non-consumed duplicate
portion: value_slice[dup_start:consumed_from]."""
return prefix_len
def should_skip_leaf_creation(
self, total_prefix_len: int, key_len: int, params: InsertParams
) -> bool:
"""Return True to veto leaf creation when the entire new leaf would
be a tombstone for this component."""
return False
def commit_insert_component_data(
self,
node: UnifiedTreeNode,
is_new_leaf: bool,
params: InsertParams,
result: InsertResult,
) -> None:
"""Finalize component data on the target (leaf) node after the insert
walk completes. Called once per insert.
- Full: no-op (full data is handled by _add_new_node).
- SWA: for new leaves, checks whether the node straddles the SWA
eviction boundary (swa_evicted_seqlen). If so, splits the node
via _split_node — the parent becomes a tombstone (no SWA) and the
child (the deeper portion) receives SWA data. If the entire node
is within the window, sets SWA directly. If entirely outside,
leaves SWA as None (tombstone).
- Mamba: sets the mamba component value from params, inserts into
mamba LRU list, and increments evictable size. If the node already
has mamba data, resets its LRU position instead."""
pass
@abstractmethod
def redistribute_on_node_split(
self, new_parent: UnifiedTreeNode, child: UnifiedTreeNode
):
"""Redistribute component data between new_parent and child when a
node is split. new_parent is the newly created prefix node.
- Full: copies child's lock_ref to new_parent.
- SWA: slices (or clones) the swa value for new_parent, copies
lock_ref and component_uuid metadata, then syncs child's swa
value with its (now-trimmed) full_value.
- Mamba: sets new_parent's mamba value to None and lock_ref to 0
(mamba data stays on the original leaf, not on prefix nodes)."""
...
@abstractmethod
def evict_component(self, node: UnifiedTreeNode, is_leaf: bool) -> int:
"""Free this component's KV resources on a node being evicted.
For internal (non-leaf) nodes: free memory and tombstone the value
(set to None); the node structure is kept.
For leaf nodes: free memory; the node will be deleted by caller.
Returns the number of tokens/slots freed.
- Full: frees full_value via token_to_kv_pool_allocator.
- SWA: frees swa value via swa_token_to_kv_pool_allocator;
only tombstones on internal nodes.
- Mamba: frees mamba value via mamba_token_to_kv_pool_allocator;
only tombstones on internal nodes."""
...
def eviction_priority(self, is_leaf: bool) -> int:
"""Eviction priority on this node type. Higher = evicted later.
When a component is evicted, all other components with equal or
lower priority on the same node are also cascade-evicted.
Leaf: all components equal (0) — evicting any cascades to all,
because the node will be deleted.
Internal: full=2 > swa=1 > mamba=0.
Why swa > mamba: SWA data on internal nodes is *path data* —
the sliding window needs continuous SWA coverage along the path
from root to the match boundary. E.g. A->B->C->D->E where C
and E both have mamba and the window covers C->E: if C's mamba
is evicted, C's SWA must stay so E remains reachable.
Mamba data, by contrast, is only meaningful at the match
boundary node; on internal nodes it
contributes nothing to the path. So SWA is more valuable to
keep and should be evicted later.
Cascade consequences:
- Mamba evict internal: no cascade.
- SWA evict internal: cascades to Mamba. SWA gone -> SWA
validator fails -> mamba data is useless (match requires all
validators to pass).
- Full evict internal: cascades to SWA + Mamba."""
return 0
@abstractmethod
def drive_eviction(
self, params: EvictParams, tracker: dict[ComponentType, int]
) -> None:
"""Drive eviction from this component's LRU list.
Each component extracts its own request from params, walks its own
LRU, evicts, and calls cache._cascade_evict for priority cascade.
Updates the shared tracker with freed amounts for all components.
- Full: walks leaf LRU, evicts full then cascades entire leaf.
- Mamba: walks full LRU; tombstones internal nodes (with cascade
to equal-priority components like swa), cascades leaves to all."""
...
@abstractmethod
def acquire_component_lock(
self, node: UnifiedTreeNode, result: IncLockRefResult
) -> IncLockRefResult:
"""Increment lock_ref for this component, protecting nodes from
eviction. Updates evictable → protected size on first lock.
- Full: path-lock — walks from node up to root, incrementing
lock_ref on every ancestor.
- SWA: path-lock — walks upward collecting swa values until the
sliding window is filled; records a component_uuid at the
boundary for release_component_lock to know where to stop.
- Mamba: single-node lock — only increments lock_ref on the
node itself (mamba state is per-leaf, not per-path)."""
...
@abstractmethod
def release_component_lock(
self, node: UnifiedTreeNode, params: Optional[DecLockRefParams]
) -> None:
"""Decrement lock_ref for this component, un-protecting nodes.
Updates protected → evictable size when lock_ref drops to 0.
- Full: path-unlock — walks from node up to root, decrementing
lock_ref on every ancestor.
- SWA: path-unlock — walks upward, stopping at the node whose
component_uuid matches the one recorded during acquire.
- Mamba: single-node unlock — only decrements lock_ref on the
node itself."""
...
def prepare_for_caching_req(
self,
req: Req,
insert_params: InsertParams,
token_ids_len: int,
is_finished: bool,
) -> Optional[int]:
"""Prepare component-specific data before insert, fill component
fields in insert_params, return effective cache_len.
Return None for no truncation opinion (use full length);
return int >= 0 for effective cache length.
- Full: no-op, returns None.
- SWA: sets insert_params.swa_evicted_seqlen on finished; returns None.
- Mamba: prepares mamba_value (finished from ping-pong buffer,
unfinished fork from req); returns mamba_last_track_seqlen."""
return None
def cleanup_after_caching_req(
self,
req: Req,
is_finished: bool,
insert_result: Optional[InsertResult] = None,
insert_params: Optional[InsertParams] = None,
) -> None:
"""Post-cache cleanup for component-specific resources.
``is_finished`` — whether the request has finished generation.
True means the request is complete and its resources can be released;
``insert_result`` is None when insert was skipped (cache disabled
or effective_cache_len <= 0); treat as "no insert happened".
``insert_params`` is None only on the disabled path; on early-return
paths it is still provided so components can free their resources."""
pass
@@ -0,0 +1,986 @@
from __future__ import annotations
import logging
import time
from collections import defaultdict
from functools import partial
from typing import TYPE_CHECKING, Optional
import torch
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
DecLockRefParams,
DecLockRefResult,
EvictParams,
EvictResult,
IncLockRefResult,
InsertParams,
InsertResult,
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.radix_cache import (
RadixKey,
_key_match_page_size1,
_key_match_paged,
get_child_key,
maybe_bigram_convert,
page_align_keys,
)
from sglang.srt.mem_cache.unified_cache_components import (
_NUM_COMPONENT_TYPES,
BASE_COMPONENT_TYPE,
ComponentData,
ComponentType,
FullComponent,
MambaComponent,
SWAComponent,
TreeComponent,
get_and_increase_time_counter,
)
from sglang.srt.mem_cache.utils import convert_to_bigram_key
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
class UnifiedTreeNode:
counter = 0
def __init__(self, tree_components: tuple[ComponentType, ...]):
self.children = defaultdict(partial(UnifiedTreeNode, tree_components))
self.parent: UnifiedTreeNode | None = None
self.key: Optional[RadixKey] = None
self.tree_components = tree_components
# list indexed by ComponentType (int enum 0..N-1)
self.component_data: list[ComponentData] = [
ComponentData() for _ in range(_NUM_COMPONENT_TYPES)
]
self.last_access_time = get_and_increase_time_counter()
self.host_value = None
self.hit_count = 0
self.lru_prev: list[UnifiedTreeNode | None] = [None] * _NUM_COMPONENT_TYPES
self.lru_next: list[UnifiedTreeNode | None] = [None] * _NUM_COMPONENT_TYPES
self.id = UnifiedTreeNode.counter
UnifiedTreeNode.counter += 1
def component(self, component_type: ComponentType) -> ComponentData:
return self.component_data[component_type]
def __lt__(self, other: UnifiedTreeNode):
return self.last_access_time < other.last_access_time
class UnifiedLRUList:
def __init__(
self, component_type: ComponentType, tree_components: tuple[ComponentType, ...]
):
self.component_type = component_type
self.head = UnifiedTreeNode(tree_components)
self.tail = UnifiedTreeNode(tree_components)
self.head.lru_next[component_type] = self.tail
self.tail.lru_prev[component_type] = self.head
self.cache: dict[int, UnifiedTreeNode] = {}
def _add_node_after(self, prev_node: UnifiedTreeNode, new_node: UnifiedTreeNode):
ct = self.component_type
new_node.lru_prev[ct] = prev_node
new_node.lru_next[ct] = prev_node.lru_next[ct]
prev_node.lru_next[ct].lru_prev[ct] = new_node
prev_node.lru_next[ct] = new_node
def _add_node(self, node: UnifiedTreeNode):
self._add_node_after(self.head, node)
def _remove_node(self, node: UnifiedTreeNode):
ct = self.component_type
node.lru_prev[ct].lru_next[ct] = node.lru_next[ct]
node.lru_next[ct].lru_prev[ct] = node.lru_prev[ct]
def insert_mru(self, node: UnifiedTreeNode):
assert node.id not in self.cache
self.cache[node.id] = node
self._add_node(node)
def remove_node(self, node: UnifiedTreeNode):
assert node.id in self.cache
del self.cache[node.id]
self._remove_node(node)
def reset_node_mru(self, node: UnifiedTreeNode):
assert node.id in self.cache
self._remove_node(node)
self._add_node(node)
def reset_node_and_parents_mru(
self,
node: UnifiedTreeNode,
root_node: UnifiedTreeNode,
should_include,
):
prev_node = self.head
while node != root_node:
if should_include(node):
assert node.id in self.cache
self._remove_node(node)
self._add_node_after(prev_node, node)
prev_node = node
node = node.parent
def in_list(self, node: Optional[UnifiedTreeNode]):
return node is not None and node.id in self.cache
def get_prev_no_lock(self, node: UnifiedTreeNode, check_id: bool = True):
if check_id:
assert node.id in self.cache
ct = self.component_type
x = node.lru_prev[ct]
while x.component_data[ct].lock_ref > 0:
x = x.lru_prev[ct]
if x == self.head:
return None
return x
def get_prev_leaf_no_lock(self, node: UnifiedTreeNode, check_id: bool = True):
if check_id:
assert node.id in self.cache
ct = self.component_type
x = node.lru_prev[ct]
while x.component_data[ct].lock_ref > 0 or len(x.children) > 0:
x = x.lru_prev[ct]
if x == self.head:
return None
return x
def get_lru_no_lock(self):
return self.get_prev_no_lock(self.tail, check_id=False)
def get_leaf_lru_no_lock(self):
return self.get_prev_leaf_no_lock(self.tail, check_id=False)
COMPONENT_REGISTRY: dict[ComponentType, type[TreeComponent]] = {
ComponentType.FULL: FullComponent,
ComponentType.MAMBA: MambaComponent,
ComponentType.SWA: SWAComponent,
}
logger = logging.getLogger(__name__)
class UnifiedRadixCache(BasePrefixCache):
def __init__(
self,
params: CacheInitParams,
):
self.req_to_token_pool = params.req_to_token_pool
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
self.page_size = params.page_size
self.disable = params.disable
self.is_eagle = params.is_eagle
if self.token_to_kv_pool_allocator:
self.device = self.token_to_kv_pool_allocator.device
else:
self.device = torch.device("cpu")
if params.enable_metrics:
self.init_metrics_collector()
if self.page_size == 1:
self.key_match_fn = _key_match_page_size1
self.get_child_key_fn = get_child_key
else:
self.key_match_fn = partial(_key_match_paged, page_size=self.page_size)
self.get_child_key_fn = partial(get_child_key, page_size=self.page_size)
assert params.tree_components is not None
self.tree_components = tuple(params.tree_components)
self.components: dict[ComponentType, TreeComponent] = {
ct: COMPONENT_REGISTRY[ct](self, params) for ct in self.tree_components
}
self._components_tuple: tuple[TreeComponent, ...] = tuple(
self.components.values()
)
if self.is_eagle:
self.key_convert_fn = convert_to_bigram_key
else:
self.key_convert_fn = lambda key: key
self.reset()
logger.info(f"Init Unified RadixTree with components {self.tree_components}")
def reset(self) -> None:
self.root_node = UnifiedTreeNode(self.tree_components)
self.root_node.key = RadixKey([], None)
self.root_node.component_data[BASE_COMPONENT_TYPE].value = []
for ct in self.tree_components:
self.root_node.component_data[ct].lock_ref = 1
self.component_evictable_size_ = {ct: 0 for ct in self.tree_components}
self.component_protected_size_ = {ct: 0 for ct in self.tree_components}
self.lru_lists = {
ct: UnifiedLRUList(ct, self.tree_components) for ct in self.tree_components
}
def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
key = params.key
key, _ = maybe_bigram_convert(self.is_eagle, key)
if self.disable or len(key) == 0:
return MatchResult(
device_indices=torch.empty(
(0,),
dtype=torch.int64,
device=self.device,
),
last_device_node=self.root_node,
last_host_node=self.root_node,
)
if self.page_size != 1:
page_aligned_len = len(key) // self.page_size * self.page_size
key = key[:page_aligned_len]
value, last_node, best_value_len = self._match_prefix_helper(key)
return self._match_post_processor(params, value, last_node, best_value_len)
def insert(self, params: InsertParams) -> InsertResult:
if self.disable:
return InsertResult(prefix_len=0)
key = params.key
value = params.value
if value is None:
value = torch.tensor([x for x in key.token_ids], dtype=torch.int64)
key, value = maybe_bigram_convert(self.is_eagle, key, value)
result = self._insert_helper(self.root_node, key, value, params)
return result
def evict(self, params: EvictParams) -> EvictResult:
if self.disable:
return EvictResult()
start_time = time.perf_counter()
tracker = {ct: 0 for ct in self.tree_components}
for component in self._components_tuple:
component.drive_eviction(params=params, tracker=tracker)
self.update_eviction_metrics(sum(tracker.values()), start_time)
return EvictResult(
num_tokens_evicted=tracker[BASE_COMPONENT_TYPE],
swa_num_tokens_evicted=tracker.get(ComponentType.SWA, 0),
mamba_num_evicted=tracker.get(ComponentType.MAMBA, 0),
)
def inc_lock_ref(self, node: UnifiedTreeNode) -> IncLockRefResult:
if self.disable:
return IncLockRefResult()
result = IncLockRefResult()
for component in self._components_tuple:
result = component.acquire_component_lock(node=node, result=result)
return result
def dec_lock_ref(
self, node: UnifiedTreeNode, params: Optional[DecLockRefParams] = None
) -> DecLockRefResult:
if self.disable:
return DecLockRefResult()
for component in self._components_tuple:
component.release_component_lock(node=node, params=params)
# TODO: delta is not aggregated from components; no caller uses it yet.
return DecLockRefResult()
def cache_finished_req(self, req: Req, is_insert: bool = True) -> None:
kv_committed_len = req.pop_committed_kv_cache()
if self.disable:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :kv_committed_len
]
self.token_to_kv_pool_allocator.free(kv_indices)
for comp in self._components_tuple:
comp.cleanup_after_caching_req(req, is_finished=True)
return
token_ids = (req.origin_input_ids + req.output_ids)[:kv_committed_len]
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :kv_committed_len
]
result = None
insert_params = None
if is_insert:
insert_params = InsertParams(prev_prefix_len=req.cache_protected_len)
# components prepare insert data + return effective cache_len
effective_cache_len = len(token_ids)
for comp in self._components_tuple:
cl = comp.prepare_for_caching_req(
req=req,
insert_params=insert_params,
token_ids_len=len(token_ids),
is_finished=True,
)
if cl is not None:
effective_cache_len = min(effective_cache_len, cl)
# Truncate if needed
if effective_cache_len < len(token_ids):
free_start = max(effective_cache_len, req.cache_protected_len)
self.token_to_kv_pool_allocator.free(kv_indices[free_start:])
token_ids = token_ids[:effective_cache_len]
kv_indices = kv_indices[:effective_cache_len]
# Key convert + page align
keys = self.key_convert_fn(token_ids)
keys = page_align_keys(keys, self.page_size)
page_aligned_len = len(keys)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
insert_params.key = radix_key
insert_params.value = values
result = self.insert(insert_params)
# Free unaligned tail
self.token_to_kv_pool_allocator.free(kv_indices[page_aligned_len:])
else:
self.token_to_kv_pool_allocator.free(kv_indices[req.cache_protected_len :])
self.dec_lock_ref(
req.last_node,
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
)
# cleanup
for comp in self._components_tuple:
comp.cleanup_after_caching_req(
req, is_finished=True, insert_result=result, insert_params=insert_params
)
def cache_unfinished_req(self, req: Req, chunked=False) -> None:
token_ids = req.fill_ids
if self.disable:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : len(token_ids)
]
req.prefix_indices = kv_indices
return
kv_indices_orig = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : len(token_ids)
]
# components prepare insert data + return effective cache_len
insert_params = InsertParams(prev_prefix_len=req.cache_protected_len)
effective_cache_len = len(token_ids)
for comp in self._components_tuple:
cl = comp.prepare_for_caching_req(
req=req,
insert_params=insert_params,
token_ids_len=len(token_ids),
is_finished=False,
)
if cl is not None:
effective_cache_len = min(effective_cache_len, cl)
if effective_cache_len <= 0:
req.prefix_indices = kv_indices_orig.to(dtype=torch.int64, copy=True)
for comp in self._components_tuple:
comp.cleanup_after_caching_req(
req, is_finished=False, insert_params=insert_params
)
return
kv_indices = kv_indices_orig[:effective_cache_len]
# Key convert + page align
keys = self.key_convert_fn(token_ids[:effective_cache_len])
keys = page_align_keys(keys, self.page_size)
page_aligned_len = len(keys)
values = kv_indices[:page_aligned_len].to(dtype=torch.int64, copy=True)
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
insert_params.key = radix_key
insert_params.value = values
result = self.insert(insert_params)
# Match prefix
match_result = self.match_prefix(MatchPrefixParams(key=radix_key))
new_indices = match_result.device_indices
new_last_node = match_result.last_device_node
new_prefix_len = result.prefix_len
assert (
req.cache_protected_len <= len(new_indices) + self.page_size - 1
), f"{req.cache_protected_len=}, {len(new_indices)=}, {page_aligned_len=}"
assert new_prefix_len <= len(
new_indices
), f"{new_prefix_len=}, {len(new_indices)=}"
self.req_to_token_pool.write(
(req.req_pool_idx, slice(req.cache_protected_len, len(new_indices))),
new_indices[req.cache_protected_len :],
)
self.dec_lock_ref(
req.last_node,
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
)
lock_result = self.inc_lock_ref(new_last_node)
# Update req fields
if len(new_indices) < len(kv_indices_orig):
req.prefix_indices = torch.cat(
[new_indices, kv_indices_orig[len(new_indices) :]]
)
else:
req.prefix_indices = new_indices
req.cache_protected_len = len(new_indices)
req.last_node = new_last_node
req.swa_uuid_for_lock = lock_result.swa_uuid_for_lock
# cleanup
for comp in self._components_tuple:
comp.cleanup_after_caching_req(
req,
is_finished=False,
insert_result=result,
insert_params=insert_params,
)
# ---- Internal Helpers ----
def _match_prefix_helper_readonly(
self, key: RadixKey
) -> tuple[list[torch.Tensor], UnifiedTreeNode, int]:
"""Read-only version of _match_prefix_helper that does not split nodes.
Only considers fully matched nodes, ignores partial matches.
Not used yet; reserved for future read-only match operations."""
node = self.root_node
child_key = self.get_child_key_fn(key)
value: list[torch.Tensor] = []
best_value_len = 0
best_node = node
validators = tuple(
comp.create_match_validator() for comp in self._components_tuple
)
def _update_best_if_valid(node):
nonlocal best_value_len, best_node
if all(v(node) for v in validators):
best_value_len = len(value)
best_node = node
while len(key) > 0 and child_key in node.children:
child = node.children[child_key]
prefix_len = self.key_match_fn(child.key, key)
if prefix_len < len(child.key):
# Read-only: do not split, ignore partial match and stop
break
value.append(child.component_data[BASE_COMPONENT_TYPE].value)
node = child
_update_best_if_valid(node)
key = key[prefix_len:]
if len(key):
child_key = self.get_child_key_fn(key)
return value, best_node, best_value_len
def _match_prefix_helper(
self, key: RadixKey
) -> tuple[list[torch.Tensor], UnifiedTreeNode, int]:
node = self.root_node
child_key = self.get_child_key_fn(key)
value: list[torch.Tensor] = []
best_value_len = 0
best_node = node
validators = tuple(
comp.create_match_validator() for comp in self._components_tuple
)
def _update_best_if_valid(node):
nonlocal best_value_len, best_node
if all(v(node) for v in validators):
best_value_len = len(value)
best_node = node
while len(key) > 0 and child_key in node.children:
child = node.children[child_key]
prefix_len = self.key_match_fn(child.key, key)
if prefix_len < len(child.key):
node = self._split_node(child.key, child, prefix_len)
value.append(node.component_data[BASE_COMPONENT_TYPE].value)
_update_best_if_valid(node)
break
value.append(child.component_data[BASE_COMPONENT_TYPE].value)
node = child
_update_best_if_valid(node)
key = key[prefix_len:]
if len(key):
child_key = self.get_child_key_fn(key)
return value, best_node, best_value_len
def _match_post_processor(
self,
params: MatchPrefixParams,
value: list[torch.Tensor],
last_node: UnifiedTreeNode,
best_value_len: int,
) -> MatchResult:
node_update = last_node
for comp in self._components_tuple:
self.lru_lists[comp.component_type].reset_node_and_parents_mru(
node_update, self.root_node, comp.node_has_component_data
)
cur_time = get_and_increase_time_counter()
while node_update:
node_update.last_access_time = cur_time
cur_time -= 0.00001
node_update = node_update.parent
if best_value_len > 0:
device_indices = torch.cat(value[:best_value_len])
else:
device_indices = torch.empty((0,), dtype=torch.int64, device=self.device)
result = MatchResult(
device_indices=device_indices,
last_device_node=last_node,
last_host_node=last_node,
)
for component in self._components_tuple:
result = component.finalize_match_result(
result=result,
params=params,
value_chunks=value,
best_value_len=best_value_len,
)
return result
def _split_node(
self, key: RadixKey, child: UnifiedTreeNode, split_len: int
) -> UnifiedTreeNode:
new_node = UnifiedTreeNode(self.tree_components)
new_node.children = {self.get_child_key_fn(key[split_len:]): child}
new_node.parent = child.parent
new_node.key = child.key[:split_len]
new_node.component_data[BASE_COMPONENT_TYPE].value = (
child.component_data[BASE_COMPONENT_TYPE].value[:split_len].clone()
)
self._for_each_component_lru(child, UnifiedLRUList.remove_node)
child.parent = new_node
child.key = child.key[split_len:]
child.component_data[BASE_COMPONENT_TYPE].value = (
child.component_data[BASE_COMPONENT_TYPE].value[split_len:].clone()
)
for component in self._components_tuple:
component.redistribute_on_node_split(new_parent=new_node, child=child)
new_node.parent.children[self.get_child_key_fn(key)] = new_node
self._for_each_component_lru(new_node, UnifiedLRUList.insert_mru)
self._for_each_component_lru(child, UnifiedLRUList.insert_mru)
child.last_access_time = get_and_increase_time_counter()
return new_node
def _touch_node(self, node: UnifiedTreeNode):
node.last_access_time = get_and_increase_time_counter()
if node != self.root_node:
self._for_each_component_lru(node, UnifiedLRUList.reset_node_mru)
def _add_new_node(
self,
parent: UnifiedTreeNode,
key: RadixKey,
value: torch.Tensor,
) -> UnifiedTreeNode:
new_node = UnifiedTreeNode(self.tree_components)
new_node.parent = parent
new_node.key = key
new_node.component_data[BASE_COMPONENT_TYPE].value = value.clone()
parent.children[self.get_child_key_fn(key)] = new_node
self.lru_lists[BASE_COMPONENT_TYPE].insert_mru(new_node)
self.component_evictable_size_[BASE_COMPONENT_TYPE] += len(value)
return new_node
def _insert_helper(
self,
node: UnifiedTreeNode,
key: RadixKey,
value: torch.Tensor,
params: InsertParams,
) -> InsertResult:
self._touch_node(node)
if len(key) == 0:
return InsertResult(prefix_len=0, mamba_exist=True)
child_key = self.get_child_key_fn(key)
total_prefix_length = 0
while len(key) > 0 and child_key in node.children:
node = node.children[child_key]
self._touch_node(node)
prefix_len = self.key_match_fn(node.key, key)
if prefix_len < len(node.key):
node = self._split_node(node.key, node, prefix_len)
value_slice = value[:prefix_len]
consumed_from = prefix_len
# Let each component claim ownership of overlapping KV slots
for component in self._components_tuple:
comp_consumed_from = component.update_component_on_insert_overlap(
node=node,
prefix_len=prefix_len,
total_prefix_len=total_prefix_length,
value_slice=value_slice,
params=params,
)
consumed_from = min(consumed_from, comp_consumed_from)
dup_start = max(0, params.prev_prefix_len - total_prefix_length)
if dup_start < consumed_from:
self.token_to_kv_pool_allocator.free(
value_slice[dup_start:consumed_from]
)
total_prefix_length += prefix_len
key = key[prefix_len:]
value = value[prefix_len:]
if len(key):
child_key = self.get_child_key_fn(key)
is_new_leaf = False
# Create new leaf for remaining suffix
if len(key):
if any(
comp.should_skip_leaf_creation(
total_prefix_len=total_prefix_length,
key_len=len(key),
params=params,
)
for comp in self._components_tuple
):
# TODO: When leaf creation is skipped, We should release all component
# resources here or propagate a flag so that
# cleanup_after_caching_req can free them properly.
self.token_to_kv_pool_allocator.free(value)
return InsertResult(prefix_len=total_prefix_length)
target_node = self._add_new_node(node, key, value)
is_new_leaf = True
else:
target_node = node
# Finalize: let each component attach its data to the target node.
# e.g. Mamba attaches mamba_value to the leaf node
result = InsertResult(prefix_len=total_prefix_length)
for component in self._components_tuple:
component.commit_insert_component_data(
node=target_node,
is_new_leaf=is_new_leaf,
params=params,
result=result,
)
return result
def _cascade_evict(
self,
node: UnifiedTreeNode,
trigger: TreeComponent,
tracker: dict[ComponentType, int],
):
"""Cascade eviction from trigger to lower-or-equal priority components.
When a component evicts a node, all other components with equal or
lower eviction_priority on the same node are also evicted.
If the node is a leaf, it is removed from the tree and any
resulting tombstone ancestors are cleaned up recursively."""
is_leaf = len(node.children) == 0
trigger_priority = trigger.eviction_priority(is_leaf)
for comp in self._components_tuple:
if comp.eviction_priority(is_leaf) <= trigger_priority:
if comp is not trigger and comp.node_has_component_data(node):
assert node.component_data[comp.component_type].lock_ref == 0
self._evict_component_and_detach_lru(
node, comp, is_leaf=is_leaf, tracker=tracker
)
if is_leaf:
self._remove_leaf_from_parent(node)
self._iteratively_delete_tombstone_leaf(node, tracker)
def _remove_leaf_from_parent(self, node: UnifiedTreeNode):
key = self.get_child_key_fn(node.key)
v = node.parent.children.pop(key, None)
assert v == node
def _evict_component_and_detach_lru(
self,
node: UnifiedTreeNode,
comp: TreeComponent,
is_leaf: bool,
tracker: dict[ComponentType, int],
) -> int:
freed = comp.evict_component(node, is_leaf=is_leaf)
tracker[comp.component_type] += freed
lru = self.lru_lists[comp.component_type]
if lru.in_list(node):
lru.remove_node(node)
return freed
def _iteratively_delete_tombstone_leaf(
self, deleted_node: UnifiedTreeNode, tracker: dict[ComponentType, int]
):
"""After a leaf is removed, walk up the parent chain and delete
any ancestor that is leaf node and has lost any component data (tombstoned)."""
cur = deleted_node.parent
while cur != self.root_node and len(cur.children) == 0:
has_tombstone = any(
not comp.node_has_component_data(cur)
for comp in self.components.values()
)
if not has_tombstone:
break
if any(
cur.component_data[comp.component_type].lock_ref > 0
for comp in self.components.values()
if comp.node_has_component_data(cur)
):
break
for comp in self.components.values():
if comp.node_has_component_data(cur):
self._evict_component_and_detach_lru(
cur, comp, is_leaf=True, tracker=tracker
)
self._remove_leaf_from_parent(cur)
cur = cur.parent
def _for_each_component_lru(self, node: UnifiedTreeNode, lru_op):
for ct in self.tree_components:
if node.component_data[ct].value is not None:
lru_op(self.lru_lists[ct], node)
# ---- Query / Inspection APIs ----
# These APIs exist for compatibility with other RadixTree implementations.
# TODO: simplify and consolidate in a future refactor.
@property
def sliding_window_size(self):
swa = self.components.get(ComponentType.SWA)
return swa.sliding_window_size if swa else None
def supports_swa(self) -> bool:
return ComponentType.SWA in self.components
def supports_mamba(self) -> bool:
return ComponentType.MAMBA in self.components
def evictable_size(self) -> int:
return self.component_evictable_size_.get(BASE_COMPONENT_TYPE, 0)
def protected_size(self) -> int:
return self.component_protected_size_.get(BASE_COMPONENT_TYPE, 0)
def full_evictable_size(self) -> int:
return self.evictable_size()
def full_protected_size(self) -> int:
return self.protected_size()
def swa_evictable_size(self) -> int:
return self.component_evictable_size_.get(ComponentType.SWA, 0)
def mamba_evictable_size(self) -> int:
return self.component_evictable_size_.get(ComponentType.MAMBA, 0)
def swa_protected_size(self) -> int:
return self.component_protected_size_.get(ComponentType.SWA, 0)
def mamba_protected_size(self) -> int:
return self.component_protected_size_.get(ComponentType.MAMBA, 0)
def total_size(self):
total_size = 0
total_aux_size = 0
stack = [self.root_node]
while stack:
node = stack.pop()
total_size += len(node.component_data[BASE_COMPONENT_TYPE].value)
for ct in self.tree_components:
if ct == BASE_COMPONENT_TYPE:
continue
value = node.component_data[ct].value
if value is not None:
total_aux_size += len(value)
for child in node.children.values():
stack.append(child)
return total_size, total_aux_size
def all_values_flatten(self) -> torch.Tensor:
values = []
def _dfs(node: UnifiedTreeNode):
for child in node.children.values():
values.append(child.component_data[BASE_COMPONENT_TYPE].value)
_dfs(child)
_dfs(self.root_node)
if values:
return torch.cat(values)
return torch.tensor([], dtype=torch.int64, device=self.device)
def _all_component_values_flatten(
self, component_type: ComponentType
) -> torch.Tensor:
if component_type not in self.components:
return torch.tensor([], dtype=torch.int64, device=self.device)
values = []
def _dfs(node: UnifiedTreeNode):
value = node.component_data[component_type].value
if value is not None:
values.append(value)
for child in node.children.values():
_dfs(child)
_dfs(self.root_node)
if values:
return torch.cat(values)
return torch.tensor([], dtype=torch.int64, device=self.device)
def all_mamba_values_flatten(self) -> torch.Tensor:
return self._all_component_values_flatten(ComponentType.MAMBA)
def all_swa_values_flatten(self) -> torch.Tensor:
return self._all_component_values_flatten(ComponentType.SWA)
def available_and_evictable_str(self) -> str:
if self.supports_swa():
full_available_size = self.token_to_kv_pool_allocator.full_available_size()
else:
full_available_size = self.token_to_kv_pool_allocator.available_size()
full_evictable = self.component_evictable_size_[BASE_COMPONENT_TYPE]
lines = [
f"Available full tokens: {full_available_size + full_evictable} "
f"(full_available_size={full_available_size} + full_evictable_size_={full_evictable})"
]
for ct in self.tree_components:
if ct == BASE_COMPONENT_TYPE:
continue
if ct.is_swa:
available_size = self.token_to_kv_pool_allocator.swa_available_size()
elif ct.is_mamba:
available_size = self.req_to_token_pool.mamba_pool.available_size()
else:
continue
lines.append(
f"Available {ct}: {available_size + self.component_evictable_size_[ct]} "
f"(available_size={available_size} + component_evictable_size_={self.component_evictable_size_[ct]})"
)
return "\n".join(lines) + "\n"
def _collect_all_nodes(self) -> list[UnifiedTreeNode]:
nodes = []
stack = [self.root_node]
while stack:
node = stack.pop()
nodes.append(node)
stack.extend(node.children.values())
return nodes
def sanity_check(self):
"""Thorough sanity check: verify LRU membership, lock state, linked-list
integrity, and evictable sizes for every component.
Expensive — use only in tests or idle checks."""
try:
# 1. Collect all nodes from tree
all_nodes = self._collect_all_nodes()
for ct in self.tree_components:
# 2. Basic size invariants
assert (
self.component_evictable_size_[ct] >= 0
), f"component_evictable_size_[{ct}] = {self.component_evictable_size_[ct]} < 0"
assert (
self.component_protected_size_[ct] >= 0
), f"component_protected_size_[{ct}] = {self.component_protected_size_[ct]} < 0"
# 3. Verify LRU membership: tree nodes with data == LRU cache entries
lru = self.lru_lists[ct]
tree_ids = {
n.id
for n in all_nodes
if n != self.root_node and n.component_data[ct].value is not None
}
lru_ids = set(lru.cache.keys())
assert tree_ids == lru_ids, (
f"[{ct}] LRU membership mismatch: "
f"in_tree_not_lru={tree_ids - lru_ids}, "
f"in_lru_not_tree={lru_ids - tree_ids}"
)
# 4. Walk LRU doubly-linked list: verify structural integrity
# and that all nodes are unlocked (idle check)
visited = set()
x = lru.head.lru_next[ct]
prev = lru.head
while x != lru.tail:
assert (
x.lru_prev[ct] == prev
), f"[{ct}] broken prev link at node {x.id}"
assert (
x.id in lru.cache
), f"[{ct}] node {x.id} in linked list but not in cache dict"
assert x.id not in visited, f"[{ct}] cycle detected at node {x.id}"
assert x.component_data[ct].lock_ref == 0, (
f"[{ct}] node {x.id} should not be locked when idle, "
f"lock_ref={x.component_data[ct].lock_ref}"
)
visited.add(x.id)
prev = x
x = x.lru_next[ct]
assert len(visited) == len(lru.cache), (
f"[{ct}] linked list has {len(visited)} nodes, "
f"cache dict has {len(lru.cache)}"
)
# 5. Verify evictable size by walking unlocked LRU nodes
recomputed = 0
x = lru.get_lru_no_lock()
while lru.in_list(x):
v = x.component_data[ct].value
recomputed += len(v) if v is not None else 0
x = lru.get_prev_no_lock(x)
assert self.component_evictable_size_[ct] == recomputed, (
f"[{ct}] evictable_size_={self.component_evictable_size_[ct]} "
f"!= recomputed={recomputed}"
)
except Exception as e:
logger.error(f"Unified RadixTree sanity check failed: {e}")
self.pretty_print()
raise
def pretty_print(self) -> None:
stack = [(self.root_node, 0)]
while stack:
node, indent = stack.pop()
component_str = " ".join(
f"{ct}={'yes' if node.component_data[ct].value is not None else 'no'}"
for ct in self.tree_components
)
print(
" " * indent,
f"[{node.id}]",
len(node.key),
f"full_lock={node.component_data[BASE_COMPONENT_TYPE].lock_ref}",
component_str,
)
for child in node.children.values():
stack.append((child, indent + 2))
+460
View File
@@ -0,0 +1,460 @@
"""Enhanced multi-turn KL divergence test helpers."""
from __future__ import annotations
from typing import Callable
from sglang.test.kl_test_utils import (
_extract_output_logprobs,
_flush_cache,
_generate,
_get_input_logprobs,
compare_kl_divergence,
get_input_ids,
)
__all__ = [
# Cache assertion callbacks
"default_prefill_cache_assert",
"default_decode_cache_assert",
"make_mamba_prefill_assert",
"make_mamba_decode_assert",
# Enhanced test helpers
"test_input_output_logprobs_match_helper",
"test_input_output_logprobs_match_prefill_cache_hit_helper",
"test_input_output_logprobs_match_decode_cache_hit_helper",
# Internal helpers (for custom inline tests)
"_replay_and_compare_kl",
# Re-exports from kl_test_utils
"get_input_ids",
"_generate",
"_flush_cache",
"_extract_output_logprobs",
]
# =============================================================================
# Cache assertion callbacks
# =============================================================================
# Prefill signature: (result, prefix_len, label) -> None
# Decode signature: (result, history_len, output_len, label) -> None
def default_prefill_cache_assert(result: dict, prefix_len: int, label: str):
"""Standard radix cache: cached_tokens == prefix_len."""
actual = result["meta_info"]["cached_tokens"]
assert (
actual == prefix_len
), f"{label}: expected cached_tokens={prefix_len}, got {actual}"
def default_decode_cache_assert(
result: dict, history_len: int, output_len: int, label: str
):
"""Standard radix cache: cached_tokens == history_len + output_len."""
expected = history_len + output_len
actual = result["meta_info"]["cached_tokens"]
assert (
actual == expected
), f"{label}: expected cached_tokens={expected}, got {actual}"
def make_mamba_prefill_assert(chunk_size: int = 64) -> Callable:
"""Mamba: cached_tokens in [rounded_down - chunk_size, rounded_down]."""
def _check(result: dict, prefix_len: int, label: str):
actual = result["meta_info"]["cached_tokens"]
upper = (prefix_len // chunk_size) * chunk_size
lower = max(0, upper - chunk_size)
assert (
lower <= actual <= upper
), f"{label}: expected cached_tokens in [{lower}, {upper}], got {actual}"
return _check
def make_mamba_decode_assert(track_interval: int = 16) -> Callable:
"""Mamba: cached_tokens = floor((history+output-1)/interval)*interval."""
def _check(result: dict, history_len: int, output_len: int, label: str):
actual = result["meta_info"]["cached_tokens"]
if output_len <= 0:
expected = history_len
else:
expected = (
(history_len + output_len - 1) // track_interval
) * track_interval
assert (
actual == expected
), f"{label}: expected cached_tokens={expected}, got {actual}"
return _check
# =============================================================================
# Internal helpers
# =============================================================================
def _replay_and_compare_kl(
base_url: str,
model_name: str,
kl_threshold: float,
replay_input_ids: list[list[int]],
output_logprobs: list[list[float]],
label: str,
batch_size: int = 1,
):
"""Flush cache, run replay prefill in batches, compare KL divergence."""
all_input_logprobs = []
for start in range(0, len(replay_input_ids), batch_size):
end = start + batch_size
all_input_logprobs.extend(
_get_input_logprobs(
base_url,
replay_input_ids[start:end],
output_logprobs[start:end],
)
)
acc = {model_name: {"kl_div": kl_threshold}}
compare_kl_divergence(all_input_logprobs, output_logprobs, acc, model_name, label)
def _interleave_order(n: int, branches_per_group: int) -> list[int] | None:
"""Build interleaved submission order for branch stress testing.
Given n items grouped into groups of branches_per_group, returns indices
that interleave branches across groups: [g0b0, g1b0, ..., g0b1, g1b1, ...].
Returns None if no interleaving is needed.
"""
if branches_per_group <= 0 or branches_per_group >= n:
return None
num_groups = n // branches_per_group
order = [
g * branches_per_group + b
for b in range(branches_per_group)
for g in range(num_groups)
]
# Append remainder indices not covered by complete groups
for i in range(num_groups * branches_per_group, n):
order.append(i)
return order
def _generate_maybe_interleaved(base_url, inputs, max_new_tokens, order=None):
"""Generate with optional interleaved submission order.
Submits inputs reordered by ``order``, then maps results back to the
original order so the caller always sees results[i] corresponds to
inputs[i].
"""
if order is None:
return _generate(base_url, inputs, max_new_tokens, return_logprob=True)
ordered = [inputs[i] for i in order]
results = _generate(base_url, ordered, max_new_tokens, return_logprob=True)
unordered = [None] * len(results)
for idx, orig in enumerate(order):
unordered[orig] = results[idx]
return unordered
# =============================================================================
# Helper 1: test_input_output_logprobs_match_helper
# =============================================================================
def test_input_output_logprobs_match_helper(
base_url: str,
model_name: str,
kl_threshold: float,
input_ids: list[list[int]],
*,
label: str = "logprobs_match",
max_new_tokens: int = 256,
# --- Multi-turn ---
# turn_suffixes[t][i] = suffix tokens for sample i at turn t+1
turn_suffixes: list[list[list[int]]] | None = None,
# --- Cache assertion (for turns > 0) ---
assert_decode_cached_tokens: Callable | None = None,
replay_batch_size: int = 1,
):
"""Verify decode logprobs match prefill replay.
Single-turn (turn_suffixes=None):
flush -> generate(input_ids) -> replay -> KL
Multi-turn (turn_suffixes provided):
flush -> generate turn 0 ->
for t in range(len(turn_suffixes)):
input = accumulated + output + suffix[t] -> generate ->
assert_decode_cached_tokens (optional) ->
replay last turn -> KL
Multi-branch: caller passes input_ids where multiple entries share
a prefix.
"""
n = len(input_ids)
num_turns = 1 + (len(turn_suffixes) if turn_suffixes else 0)
print(f"[{label}] {n} samples, {num_turns} turns, max_new_tokens={max_new_tokens}")
_flush_cache(base_url)
current_input = list(input_ids)
last_outputs = None
prev_input_lens = [0] * n
prev_output_lens = [0] * n
for turn in range(num_turns):
if turn > 0:
suffixes = turn_suffixes[turn - 1]
current_input = [
current_input[i] + last_outputs[i] + suffixes[i] for i in range(n)
]
results = _generate(
base_url, current_input, max_new_tokens, return_logprob=True
)
assert len(results) == n
if turn > 0 and assert_decode_cached_tokens:
for i, result in enumerate(results):
assert_decode_cached_tokens(
result,
prev_input_lens[i],
prev_output_lens[i],
f"{label}[turn{turn}][{i}]",
)
last_outputs = [r["output_ids"] for r in results]
prev_input_lens = [len(current_input[i]) for i in range(n)]
prev_output_lens = [len(last_outputs[i]) for i in range(n)]
# Replay last turn
replay_ids = [current_input[i] + results[i]["output_ids"] for i in range(n)]
output_lps = [_extract_output_logprobs(r) for r in results]
_replay_and_compare_kl(
base_url,
model_name,
kl_threshold,
replay_ids,
output_lps,
label=label,
batch_size=replay_batch_size,
)
# =============================================================================
# Helper 2: test_input_output_logprobs_match_prefill_cache_hit_helper
# =============================================================================
def test_input_output_logprobs_match_prefill_cache_hit_helper(
base_url: str,
model_name: str,
kl_threshold: float,
input_ids: list[list[int]] | None = None,
*,
# --- Multi-branch: explicit prefix/full split ---
prefix_input_ids: list[list[int]] | None = None,
full_input_ids: list[list[int]] | None = None,
label: str = "prefill_cache_hit",
max_new_tokens: int = 256,
# --- Multi-turn: additional turns after the cache-hit generation ---
turn_suffixes: list[list[list[int]]] | None = None,
# --- Cache assertions ---
assert_prefill_cached_tokens: Callable | None = None, # turn 0
assert_decode_cached_tokens: Callable | None = None, # turns > 0
# --- Interleaving for branch stress ---
branches_per_group: int = 0,
replay_batch_size: int = 1,
):
"""Verify logprobs when prefill cache is hit.
Original (input_ids only, backward compat):
flush -> seed(input_ids) -> generate(input_ids, cache hit) -> replay -> KL
Multi-branch (prefix_input_ids + full_input_ids):
flush -> seed(prefixes) -> generate(fulls, prefix cache hit) ->
assert_prefill_cached_tokens -> replay -> KL
Multi-turn (+ turn_suffixes):
... after prefill cache-hit turn, additional turns:
input = accumulated + output + suffix -> generate ->
assert_decode_cached_tokens -> replay last turn -> KL
Interleaving (branches_per_group > 0):
Reorders submission for decode-cache-hit turns to interleave branches
across groups, stressing the radix tree with competing branches.
"""
# Resolve inputs: backward compat with input_ids-only
if input_ids is not None and prefix_input_ids is None:
prefix_input_ids = input_ids
full_input_ids = input_ids
assert prefix_input_ids is not None and full_input_ids is not None
assert len(prefix_input_ids) == len(full_input_ids)
if assert_prefill_cached_tokens is None:
assert_prefill_cached_tokens = default_prefill_cache_assert
n = len(full_input_ids)
num_turns = 1 + (len(turn_suffixes) if turn_suffixes else 0)
order = _interleave_order(n, branches_per_group)
print(f"[{label}] {n} samples, {num_turns} turns, max_new_tokens={max_new_tokens}")
# Seed cache with prefixes
_flush_cache(base_url)
_generate(base_url, prefix_input_ids, max_new_tokens=0)
# Turn 0: prefill cache hit (NOT interleaved, matching original behavior)
results = _generate(base_url, full_input_ids, max_new_tokens, return_logprob=True)
assert len(results) == n
for i, result in enumerate(results):
assert_prefill_cached_tokens(
result, len(prefix_input_ids[i]), f"{label}[turn0][{i}]"
)
current_input = list(full_input_ids)
last_outputs = [r["output_ids"] for r in results]
prev_input_lens = [len(full_input_ids[i]) for i in range(n)]
prev_output_lens = [len(last_outputs[i]) for i in range(n)]
# Additional turns: decode cache hits (interleaved if order is set)
if turn_suffixes:
if assert_decode_cached_tokens is None:
assert_decode_cached_tokens = default_decode_cache_assert
for t, suffixes in enumerate(turn_suffixes):
current_input = [
current_input[i] + last_outputs[i] + suffixes[i] for i in range(n)
]
results = _generate_maybe_interleaved(
base_url, current_input, max_new_tokens, order
)
assert len(results) == n
for i, result in enumerate(results):
assert_decode_cached_tokens(
result,
prev_input_lens[i],
prev_output_lens[i],
f"{label}[turn{t + 1}][{i}]",
)
last_outputs = [r["output_ids"] for r in results]
prev_input_lens = [len(current_input[i]) for i in range(n)]
prev_output_lens = [len(last_outputs[i]) for i in range(n)]
# Replay last turn
replay_ids = [current_input[i] + results[i]["output_ids"] for i in range(n)]
output_lps = [_extract_output_logprobs(r) for r in results]
_replay_and_compare_kl(
base_url,
model_name,
kl_threshold,
replay_ids,
output_lps,
label=label,
batch_size=replay_batch_size,
)
# =============================================================================
# Helper 3: test_input_output_logprobs_match_decode_cache_hit_helper
# =============================================================================
def test_input_output_logprobs_match_decode_cache_hit_helper(
base_url: str,
model_name: str,
kl_threshold: float,
first_turn_input_ids: list[list[int]],
*,
# --- Multi-turn ---
# turn_suffixes[t][i] = suffix for sample i at turn t+2
turn_suffixes: list[list[list[int]]],
label: str = "decode_cache_hit",
max_new_tokens: int = 256,
# --- Cache assertion ---
assert_decode_cached_tokens: Callable | None = None,
# --- Interleaving ---
branches_per_group: int = 0,
replay_batch_size: int = 1,
):
"""Verify logprobs when decode cache is hit.
2-turn (turn_suffixes has 1 entry):
flush -> generate turn 1 ->
turn 2: input = turn1 + output + suffix -> generate ->
assert_decode_cached_tokens -> replay -> KL
Multi-turn (turn_suffixes has N entries):
flush -> generate turn 1 ->
for each turn t: input = accumulated + output + suffix[t] -> generate ->
assert_decode_cached_tokens -> replay last turn -> KL
Multi-branch: caller duplicates first_turn_input_ids entries and provides
different suffixes per branch. Use branches_per_group for interleaved
submission to stress the radix tree.
"""
assert (
len(turn_suffixes) >= 1
), "turn_suffixes must have at least 1 entry (for turn 2)"
if assert_decode_cached_tokens is None:
assert_decode_cached_tokens = default_decode_cache_assert
n = len(first_turn_input_ids)
num_turns = 1 + len(turn_suffixes)
order = _interleave_order(n, branches_per_group)
print(f"[{label}] {n} samples, {num_turns} turns, max_new_tokens={max_new_tokens}")
# Turn 1: populate cache, no assertion, no interleaving
_flush_cache(base_url)
results = _generate(
base_url, first_turn_input_ids, max_new_tokens, return_logprob=True
)
assert len(results) == n
current_input = list(first_turn_input_ids)
last_outputs = [r["output_ids"] for r in results]
prev_input_lens = [len(first_turn_input_ids[i]) for i in range(n)]
prev_output_lens = [len(last_outputs[i]) for i in range(n)]
# Turns 2..N: decode cache hits (interleaved if order is set)
for t, suffixes in enumerate(turn_suffixes):
current_input = [
current_input[i] + last_outputs[i] + suffixes[i] for i in range(n)
]
results = _generate_maybe_interleaved(
base_url, current_input, max_new_tokens, order
)
assert len(results) == n
for i, result in enumerate(results):
assert_decode_cached_tokens(
result,
prev_input_lens[i],
prev_output_lens[i],
f"{label}[turn{t + 1}][{i}]",
)
last_outputs = [r["output_ids"] for r in results]
prev_input_lens = [len(current_input[i]) for i in range(n)]
prev_output_lens = [len(last_outputs[i]) for i in range(n)]
# Replay last turn
replay_ids = [current_input[i] + results[i]["output_ids"] for i in range(n)]
output_lps = [_extract_output_logprobs(r) for r in results]
_replay_and_compare_kl(
base_url,
model_name,
kl_threshold,
replay_ids,
output_lps,
label=label,
batch_size=replay_batch_size,
)
@@ -0,0 +1,266 @@
import random
import unittest
from types import SimpleNamespace
from urllib.parse import urlparse
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kl_multiturn_utils import (
get_input_ids,
make_mamba_decode_assert,
make_mamba_prefill_assert,
test_input_output_logprobs_match_decode_cache_hit_helper,
test_input_output_logprobs_match_helper,
test_input_output_logprobs_match_prefill_cache_hit_helper,
)
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
def _random_suffixes(n, length, seed):
"""Generate n random token-id lists of the given length."""
rng = random.Random(seed)
return [[rng.randint(1, 30000) for _ in range(length)] for _ in range(n)]
MAMBA_MODEL = "Qwen/Qwen3-Next-80B-A3B-Instruct"
MAMBA_CHUNK_SIZE = 64
MAMBA_TRACK_INTERVAL = 128
SWA_MODEL = "openai/gpt-oss-20b"
FULL_MODEL = "Qwen/Qwen3-32B"
register_cuda_ci(est_time=1200, suite="stage-c-test-4-gpu-h100")
class UnifiedRadixTreeTestMixin:
"""Mixin: gsm8k、mmlu and multi-turn KL tests with multi-branch interleaving."""
kl_threshold: float = 0.003
max_new_tokens: int = 512
num_groups: int = 3
branches_per_group: int = 3
prefix_len: int = 512
prefill_cache_assert = None
decode_cache_assert = None
gsm8k_threshold: float = 0.93
mmlu_threshold: float = 0.8
num_gsm8k_questions: int = 200
def test_gsm8k(self):
"""Few-shot GSM8K math reasoning accuracy."""
from sglang.test.few_shot_gsm8k import run_eval as run_few_shot_gsm8k
url = urlparse(self.base_url)
args = SimpleNamespace(
num_shots=10,
data_path=None,
num_questions=self.num_gsm8k_questions,
max_new_tokens=16000,
parallel=128,
host=f"http://{url.hostname}",
port=int(url.port),
)
metrics = run_few_shot_gsm8k(args)
print(
f"[{self.__class__.__name__}] GSM8K accuracy: {metrics['accuracy']:.3f} "
f"(threshold: {self.gsm8k_threshold})"
)
self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold)
def test_mmlu(self):
"""Simple-evals MMLU multi-task accuracy."""
from sglang.test.run_eval import run_eval as run_simple_eval
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="mmlu",
num_examples=64,
num_threads=32,
)
metrics = run_simple_eval(args)
print(
f"[{self.__class__.__name__}] MMLU score: {metrics['score']:.3f} "
f"(threshold: {self.mmlu_threshold})"
)
self.assertGreaterEqual(metrics["score"], self.mmlu_threshold)
def test_multiturn_logprobs_match(self):
"""Helper 1: 3-turn, no explicit cache seeding."""
ids = self.input_ids[:4]
n = len(ids)
t2 = _random_suffixes(n, 512, seed=100)
t3 = _random_suffixes(n, 256, seed=200)
test_input_output_logprobs_match_helper(
self.base_url,
self.model,
self.kl_threshold,
ids,
turn_suffixes=[t2, t3],
assert_decode_cached_tokens=self.decode_cache_assert,
max_new_tokens=self.max_new_tokens,
)
def test_multiturn_prefill_cache_hit_branching(self):
"""Helper 2: prefill hit + 2 decode-hit turns, multi-branch interleaved."""
num_groups = self.num_groups
branches = self.branches_per_group
n = num_groups * branches
rng = random.Random(456)
prefix_ids, full_ids = [], []
for g in range(num_groups):
prefix = self.input_ids[g][: self.prefix_len]
for b in range(branches):
suffix = [rng.randint(1, 30000) for _ in range(256 + b * 64)]
prefix_ids.append(list(prefix))
full_ids.append(prefix + suffix)
t2 = _random_suffixes(n, 512, seed=789)
t3 = _random_suffixes(n, 256, seed=890)
test_input_output_logprobs_match_prefill_cache_hit_helper(
self.base_url,
self.model,
self.kl_threshold,
prefix_input_ids=prefix_ids,
full_input_ids=full_ids,
turn_suffixes=[t2, t3],
assert_prefill_cached_tokens=self.prefill_cache_assert,
assert_decode_cached_tokens=self.decode_cache_assert,
branches_per_group=branches,
max_new_tokens=self.max_new_tokens,
)
def test_multiturn_decode_cache_hit_branching(self):
"""Helper 3: 3-turn decode hit, multi-branch interleaved."""
num_groups = self.num_groups
branches = self.branches_per_group
n = num_groups * branches
first_turn = []
for g in range(num_groups):
base = self.input_ids[g][: self.prefix_len]
for _ in range(branches):
first_turn.append(list(base))
t2 = _random_suffixes(n, 512, seed=300)
t3 = _random_suffixes(n, 256, seed=400)
test_input_output_logprobs_match_decode_cache_hit_helper(
self.base_url,
self.model,
self.kl_threshold,
first_turn,
turn_suffixes=[t2, t3],
assert_decode_cached_tokens=self.decode_cache_assert,
branches_per_group=branches,
max_new_tokens=self.max_new_tokens,
)
class TestUnifiedFullRadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""Full attention."""
kl_threshold = 0.0025
@classmethod
def setUpClass(cls):
cls.model = FULL_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--mem-fraction-static",
"0.80",
"--page-size",
"64",
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
class TestUnifiedMambaRadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""Mamba hybrid + UnifiedRadixCache."""
kl_threshold = 0.003
prefill_cache_assert = staticmethod(
make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE)
)
decode_cache_assert = staticmethod(
make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL)
)
@classmethod
def setUpClass(cls):
cls.model = MAMBA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--chunked-prefill-size",
"2048",
"--mem-fraction-static",
"0.85",
"--mamba-scheduler-strategy",
"extra_buffer",
"--mamba-track-interval",
str(MAMBA_TRACK_INTERVAL),
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
"""SWA hybrid + UnifiedRadixCache."""
kl_threshold = 0.03
gsm8k_threshold = 0.75
mmlu_threshold = 0.75
@classmethod
def setUpClass(cls):
cls.model = SWA_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--mem-fraction-static",
"0.7",
"--disable-piecewise-cuda-graph",
],
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "0"},
)
cls.input_ids = get_input_ids(cls.model, num_samples=18)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,775 @@
"""Large-scale benchmark + fuzz correctness tests for UnifiedRadixCache.
Usage (standalone):
bench: python3 test/registered/unit/mem_cache/test_unified_radix_cache_bench.py --num-seqs 5000 --verify --components mamba legacy-mamba swa legacy-swa
CI Test: python -m pytest test/registered/unit/mem_cache/test_unified_radix_cache_bench.py -v -s
"""
import argparse
import gc
import logging
import random
import statistics
import time
import unittest
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Callable
import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
from sglang.srt.mem_cache.unified_cache_components.tree_component import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, suite="stage-b-test-1-gpu-small")
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
_PAGE_SIZE = 1
_HEAD_NUM = 2
_HEAD_DIM = 16
_NUM_LAYERS = 8
_GLOBAL_INTERVAL = 4
_DTYPE = torch.bfloat16
_SWA_WINDOW_SIZE = 128
_BENCH_NUM_SEQS = 5000
_BENCH_KV_SIZE = 500_000
_BENCH_CHUNK_LEN = 256
_DEFAULT_COMPONENTS = (ComponentType.FULL, ComponentType.MAMBA)
@contextmanager
def _suppress_logs():
root = logging.getLogger()
prev = root.level
root.setLevel(logging.WARNING)
try:
yield
finally:
root.setLevel(prev)
def _full_attention_layer_ids():
return list(range(_GLOBAL_INTERVAL - 1, _NUM_LAYERS, _GLOBAL_INTERVAL))
def _non_full_layer_ids():
full = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full]
# ===================================================================
# Sequence generator
# ===================================================================
def gen_random_sequences(
num_seqs: int = 2000,
chunk_len: int = 256,
vocab_size: int = 32000,
seed: int = 42,
) -> list[list[int]]:
"""Generate *num_seqs* token sequences with tree-like prefix sharing.
Phase 1 (50%): chain growth — each new seq extends a random existing one.
Phase 2 (50%): fan-out burst — multiple children from the same parent.
"""
rng = random.Random(seed)
root_prefix = [rng.randint(1, vocab_size) for _ in range(max(1, chunk_len // 4))]
sequences: list[list[int]] = [root_prefix[:]]
# Phase 1: chain growth
for _ in range(num_seqs // 2):
parent = rng.choice(sequences)
sequences.append(
parent + [rng.randint(1, vocab_size)] * rng.randint(1, chunk_len)
)
# Phase 2: fan-out burst
remaining = num_seqs - num_seqs // 2
while remaining > 0:
fan = min(rng.randint(2, 10), remaining)
parent = rng.choice(sequences)
for _ in range(fan):
sequences.append(
parent + [rng.randint(1, vocab_size)] * rng.randint(1, chunk_len)
)
remaining -= fan
rng.shuffle(sequences)
return sequences
# ===================================================================
# Cache factory
# ===================================================================
def create_bench_cache(
kv_size,
max_num_reqs,
max_context_len,
components,
page_size=_PAGE_SIZE,
tree_cls=None,
):
"""Create cache. Returns (tree, allocator, req_to_token_pool, make_req)."""
device = get_device()
has_mamba = ComponentType.MAMBA in components
has_swa = ComponentType.SWA in components
mamba2_cache_params = None
if has_mamba:
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=256,
n_groups=1,
num_heads=2,
head_dim=16,
state_size=16,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_non_full_layer_ids()
)
# --- req_to_token pool ---
if has_mamba:
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=max(max_num_reqs * 2, 200),
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_non_full_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
else:
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
# --- KV pool + allocator ---
if has_swa:
from sglang.srt.mem_cache.swa_memory_pool import (
SWAKVPool,
SWATokenToKVPoolAllocator,
)
pool = SWAKVPool(
size=kv_size,
size_swa=kv_size,
page_size=page_size,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
swa_attention_layer_ids=_non_full_layer_ids(),
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
)
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size,
page_size=page_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
else:
pool = HybridLinearKVPool(
size=kv_size,
dtype=_DTYPE,
page_size=page_size,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
enable_memory_saver=False,
mamba_pool=req_to_token_pool.mamba_pool if has_mamba else None,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
# --- tree ---
if tree_cls is None:
tree_cls = UnifiedRadixCache
tree = tree_cls(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=page_size,
disable=False,
tree_components=components if tree_cls is UnifiedRadixCache else None,
sliding_window_size=_SWA_WINDOW_SIZE if has_swa else None,
)
)
_rid = [0]
def make_req():
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.sampling.sampling_params import SamplingParams
req = Req(
rid=_rid[0],
origin_input_text="",
origin_input_ids=[],
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
)
_rid[0] += 1
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ===================================================================
# Shared bench environment + helpers
# ===================================================================
@dataclass
class _Env:
tree: object
alloc: object
rtp: object
make_req: Callable
seqs: list
has_mamba: bool
avg_tokens: int
def _make_env(num_seqs, chunk_len, kv_size, components, tree_cls=None):
"""Create sequences + cache, return shared _Env."""
if components is None:
components = _DEFAULT_COMPONENTS
seqs = gen_random_sequences(num_seqs=num_seqs, chunk_len=chunk_len)
max_seq_len = max(len(s) for s in seqs)
avg_tokens = sum(len(s) for s in seqs) // len(seqs)
with _suppress_logs():
tree, alloc, rtp, make_req = create_bench_cache(
kv_size=kv_size,
max_num_reqs=num_seqs + 100,
max_context_len=max_seq_len + 10,
components=components,
tree_cls=tree_cls,
)
return _Env(
tree, alloc, rtp, make_req, seqs, ComponentType.MAMBA in components, avg_tokens
)
def _alloc_with_evict(env, n):
"""Alloc *n* tokens, evicting if necessary. Returns tensor or None."""
v = env.alloc.alloc(n)
if v is None:
env.tree.evict(EvictParams(num_tokens=n * 2, mamba_num=2))
v = env.alloc.alloc(n)
return v
def _insert_seq(env, seq):
"""Insert one sequence (alloc + evict-fallback). Returns True on success."""
v = _alloc_with_evict(env, len(seq))
if v is None:
return False
mamba_val = None
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
env.tree.insert(InsertParams(key=RadixKey(seq), value=v, mamba_value=mamba_val))
return True
def _populate(env, count):
"""Insert first *count* sequences (with evict-fallback)."""
for seq in env.seqs[:count]:
_insert_seq(env, seq)
def _fill_no_evict(env):
"""Insert sequences until pool exhausted (no eviction). Returns count."""
inserted = 0
for seq in env.seqs:
v = env.alloc.alloc(len(seq))
if v is None:
break
mamba_val = None
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
env.tree.insert(InsertParams(key=RadixKey(seq), value=v, mamba_value=mamba_val))
inserted += 1
return inserted
# ===================================================================
# Benchmark result + runner
# ===================================================================
@dataclass
class BenchResult:
name: str
num_ops: int
total_tokens: int
elapsed_s: float
latencies_us: list[float]
@property
def ops_per_sec(self):
return self.num_ops / self.elapsed_s if self.elapsed_s > 0 else 0
@property
def tokens_per_sec(self):
return self.total_tokens / self.elapsed_s if self.elapsed_s > 0 else 0
@property
def p50_us(self):
return statistics.median(self.latencies_us) if self.latencies_us else 0
@property
def p99_us(self):
if not self.latencies_us:
return 0
idx = int(len(self.latencies_us) * 0.99)
return sorted(self.latencies_us)[min(idx, len(self.latencies_us) - 1)]
def report(self):
tok = (
f"{self.tokens_per_sec:>12,.0f} tok/s"
if self.total_tokens > 0
else f"{'N/A':>12s} tok/s"
)
return (
f" {self.name:<18s} | {tok} | {self.ops_per_sec:>10,.0f} ops/s | "
f"p50={self.p50_us:>8,.0f}us p99={self.p99_us:>8,.0f}us"
)
def bench_api(
name, setup_fn, op_fn, num_ops, tokens_per_op=0, warmup=10, verify_fn=None
):
"""Time *op_fn(item)* for each item from *setup_fn()*.
*verify_fn*, if provided, runs during warmup and once after timing
(excluded from latency measurement).
"""
items = setup_fn()
assert (
len(items) >= num_ops + warmup
), f"need {num_ops + warmup} items, got {len(items)}"
for i in range(warmup):
op_fn(items[i])
if verify_fn:
verify_fn(items[i])
gc.collect()
gc_was = gc.isenabled()
gc.disable()
latencies: list[float] = []
t0 = time.perf_counter()
for i in range(warmup, warmup + num_ops):
ts = time.perf_counter()
op_fn(items[i])
latencies.append((time.perf_counter() - ts) * 1e6)
elapsed = time.perf_counter() - t0
if gc_was:
gc.enable()
if verify_fn:
verify_fn(items[warmup + num_ops - 1])
return BenchResult(
name,
num_ops,
tokens_per_op * num_ops if tokens_per_op > 0 else 0,
elapsed,
latencies,
)
# ===================================================================
# Five benchmark scenarios
# ===================================================================
def bench_insert(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Insert throughput (alloc + evict-fallback + insert)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
warmup = min(20, num_seqs // 10)
return bench_api(
"insert",
lambda: list(range(len(env.seqs))),
lambda idx: _insert_seq(env, env.seqs[idx]),
num_seqs - warmup,
env.avg_tokens,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_match_prefix(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Prefix matching throughput (hit / partial / miss mix)."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
_populate(env, num_seqs // 2)
rng = random.Random(123)
pop = num_seqs // 2
queries: list[list[int]] = []
for _ in env.seqs:
roll = rng.random()
if roll < 0.33:
queries.append(env.seqs[rng.randint(0, pop - 1)])
elif roll < 0.66:
base = env.seqs[rng.randint(0, pop - 1)]
queries.append(base + [rng.randint(1, 32000)] * rng.randint(10, 100))
else:
queries.append([rng.randint(1, 32000)] * rng.randint(50, 300))
def verify_fn(q):
r1 = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q)))
r2 = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q)))
assert len(r1.device_indices) == len(r2.device_indices), "match not idempotent"
warmup = min(20, len(queries) // 10)
return bench_api(
"match_prefix",
lambda: queries,
lambda q: env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q))),
min(len(queries) - warmup, num_seqs),
env.avg_tokens,
warmup,
verify_fn if verify else None,
)
def bench_evict(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Eviction throughput — fill pool then repeatedly evict batches."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
inserted = _fill_no_evict(env)
evict_batch = max(100, kv_size // 200)
num_evictions = max(inserted // 5, 100)
items = [(evict_batch,)] * (num_evictions + 50)
warmup = min(20, num_evictions // 10)
return bench_api(
"evict",
lambda: items,
lambda item: env.tree.evict(EvictParams(num_tokens=item[0], mamba_num=2)),
num_evictions - warmup,
evict_batch,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_lock_unlock(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""Lock/unlock throughput — match nodes then cycle lock/unlock."""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
_populate(env, num_seqs // 2)
nodes = []
for seq in env.seqs[: num_seqs // 2]:
r = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
if r.last_device_node != env.tree.root_node:
nodes.append(r.last_device_node)
if not nodes:
return BenchResult("lock_unlock", 0, 0, 0, [])
rng = random.Random(99)
num_pairs = min(len(nodes) * 2, num_seqs)
items = [rng.choice(nodes) for _ in range(num_pairs + 50)]
def op_fn(node):
lr = env.tree.inc_lock_ref(node)
env.tree.dec_lock_ref(
node,
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
)
warmup = min(20, num_pairs // 10)
return bench_api(
"lock_unlock",
lambda: items,
op_fn,
num_pairs - warmup,
0,
warmup,
(lambda _: env.tree.sanity_check()) if verify else None,
)
def bench_cache_finished(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
tree_cls=None,
):
"""cache_finished_req throughput — full request lifecycle.
Simulates: match_prefix → inc_lock_ref → alloc → fill req_to_token → cache_finished_req.
"""
env = _make_env(num_seqs, chunk_len, kv_size, components, tree_cls)
# Pre-build Req objects with token IDs filled into req_to_token
req_items: list = []
for seq in env.seqs:
key = RadixKey(seq)
mr = env.tree.match_prefix(MatchPrefixParams(key=key))
matched_len = len(mr.device_indices)
node = mr.last_device_node
lr = env.tree.inc_lock_ref(node)
remaining = len(seq) - matched_len
if remaining > 0:
v = _alloc_with_evict(env, remaining)
if v is None:
env.tree.dec_lock_ref(
node,
DecLockRefParams(
swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)
),
)
continue
kv_indices = torch.cat([mr.device_indices, v])
else:
kv_indices = mr.device_indices
req = env.make_req()
req.origin_input_ids = list(seq)
req.output_ids = []
req.fill_ids = list(seq)
req.last_node = node
req.cache_protected_len = matched_len
req.kv_committed_len = len(seq)
req.kv_committed_freed = False
if hasattr(lr, "swa_uuid_for_lock"):
req.swa_uuid_for_lock = lr.swa_uuid_for_lock
env.rtp.req_to_token[req.req_pool_idx, : len(kv_indices)] = kv_indices
req_items.append(req)
if not req_items:
return BenchResult("cache_finished", 0, 0, 0, [])
warmup = min(20, len(req_items) // 10)
return bench_api(
"cache_finished",
lambda: req_items,
lambda req: env.tree.cache_finished_req(req, is_insert=True),
len(req_items) - warmup,
env.avg_tokens,
warmup,
# Pool math doesn't hold here (many reqs still hold allocated tokens).
(lambda _: env.tree.sanity_check()) if verify else None,
)
# ===================================================================
# Runner
# ===================================================================
ALL_BENCHMARKS = {
"insert": bench_insert,
"match": bench_match_prefix,
"evict": bench_evict,
"lock": bench_lock_unlock,
"cache_finished": bench_cache_finished,
}
def run_all_benchmarks(
num_seqs=5000,
chunk_len=256,
kv_size=500_000,
components=None,
verify=False,
benchmarks=None,
tree_cls=None,
):
if components is None:
components = _DEFAULT_COMPONENTS
if benchmarks is None or "all" in benchmarks:
benchmarks = list(ALL_BENCHMARKS.keys())
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
impl_name = (tree_cls or UnifiedRadixCache).__name__
results = []
for name in benchmarks:
if name not in ALL_BENCHMARKS:
print(f"[WARN] Unknown benchmark: {name}, skipping")
continue
results.append(
ALL_BENCHMARKS[name](
num_seqs=num_seqs,
chunk_len=chunk_len,
kv_size=kv_size,
components=components,
verify=verify,
tree_cls=tree_cls,
)
)
print("=" * 100)
print(
f"{impl_name} Benchmark | "
f"num_seqs={num_seqs} chunk_len={chunk_len} kv_size={kv_size} "
f"components={[c.value for c in components]} verify={verify}"
)
print("-" * 100)
for r in results:
print(r.report())
print("=" * 100)
return results
# ===================================================================
# pytest wrapper
# ===================================================================
class TestUnifiedRadixCacheBench(unittest.TestCase):
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def test_bench_insert(self):
r = bench_insert(_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
def test_bench_match_prefix(self):
r = bench_match_prefix(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
def test_bench_evict(self):
r = bench_evict(_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True)
self.assertGreater(r.num_ops, 0)
def test_bench_lock_unlock(self):
r = bench_lock_unlock(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
def test_bench_cache_finished(self):
r = bench_cache_finished(
_BENCH_NUM_SEQS, _BENCH_CHUNK_LEN, _BENCH_KV_SIZE, verify=True
)
self.assertGreater(r.num_ops, 0)
self.assertGreater(r.ops_per_sec, 0)
# ===================================================================
# CLI
# ===================================================================
_TREE_CONFIGS = {
"full": ((ComponentType.FULL,), None),
"mamba": ((ComponentType.FULL, ComponentType.MAMBA), None),
"swa": ((ComponentType.FULL, ComponentType.SWA), None),
"all": ((ComponentType.FULL, ComponentType.SWA, ComponentType.MAMBA), None),
"legacy-mamba": ((ComponentType.FULL, ComponentType.MAMBA), MambaRadixCache),
"legacy-swa": ((ComponentType.FULL, ComponentType.SWA), SWARadixCache),
}
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="UnifiedRadixCache benchmark")
parser.add_argument("--num-seqs", type=int, default=5000)
parser.add_argument("--chunk-len", type=int, default=256)
parser.add_argument("--kv-size", type=int, default=500_000)
parser.add_argument(
"--components",
nargs="+",
choices=list(_TREE_CONFIGS.keys()),
default=["mamba", "legacy-mamba"],
help="Component configs to benchmark",
)
parser.add_argument(
"--verify", action="store_true", help="Enable correctness assertions"
)
parser.add_argument(
"--benchmarks",
nargs="+",
default=["all"],
help="insert match evict lock cache_finished all",
)
args, _ = parser.parse_known_args()
for comp_name in args.components:
components, tree_cls = _TREE_CONFIGS[comp_name]
run_all_benchmarks(
num_seqs=args.num_seqs,
chunk_len=args.chunk_len,
kv_size=args.kv_size,
components=components,
verify=args.verify,
benchmarks=args.benchmarks,
tree_cls=tree_cls,
)
@@ -0,0 +1,883 @@
"""Unit tests for UnifiedRadixCache"""
import unittest
import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
EvictResult,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import available_and_evictable_str
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.unified_cache_components.tree_component import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import (
UnifiedRadixCache,
)
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, suite="stage-b-test-1-gpu-small")
# ---------------------------------------------------------------------------
# Shared constants
# ---------------------------------------------------------------------------
_PAGE_SIZE = 1
_HEAD_NUM = 2
_HEAD_DIM = 128
_NUM_LAYERS = 24
_GLOBAL_INTERVAL = 4
_DTYPE = torch.bfloat16
def _full_attention_layer_ids():
return [i for i in range(_GLOBAL_INTERVAL - 1, _NUM_LAYERS, _GLOBAL_INTERVAL)]
def _mamba_layer_ids():
full_set = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full_set]
def _swa_attention_layer_ids():
full_set = set(_full_attention_layer_ids())
return [i for i in range(_NUM_LAYERS) if i not in full_set]
# ===================================================================
# Test: Full + Mamba components (no SWA)
# ===================================================================
class TestUnifiedRadixCacheMamba(unittest.TestCase):
"""UnifiedRadixCache with (Full, Mamba) components."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
max_num_reqs: int = 10,
mamba_cache_size: int = 20,
max_context_len: int = 128,
):
device = get_device()
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=4096,
n_groups=16,
num_heads=32,
head_dim=128,
state_size=128,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_mamba_layer_ids()
)
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_mamba_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
pool = HybridLinearKVPool(
size=kv_size,
dtype=_DTYPE,
page_size=_PAGE_SIZE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
enable_memory_saver=False,
mamba_pool=req_to_token_pool.mamba_pool,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
tree_components=(ComponentType.FULL, ComponentType.MAMBA),
),
)
def make_req():
sp = SamplingParams(temperature=0, max_new_tokens=1)
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
sampling_params=sp,
)
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ------- insert + match -------
def test_insert_and_match_basic(self):
tree, alloc, _, make_req = self._build_tree()
# Insert [1,2,3]
req1 = make_req()
v1 = alloc.alloc(3)
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=v1,
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(result.prefix_len, 0)
# Insert [1,2,3,4,5] — shares prefix [1,2,3]
req2 = make_req()
v2 = alloc.alloc(5)
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=v2,
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(result.prefix_len, 3)
# Match [1,2,3,4,5] — full hit
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 5)
# Match [1,2,3,4,5,6] — partial hit (5 tokens)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5, 6])))
self.assertEqual(len(m.device_indices), 5)
# Match [10,11] — no hit
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([10, 11])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- evict: full-only -------
def test_evict_full_tokens(self):
tree, alloc, _, make_req = self._build_tree()
# Insert two disjoint sequences
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12]),
value=alloc.alloc(3),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 6)
# Evict 3 full tokens — should remove one leaf
result = tree.evict(EvictParams(num_tokens=3))
self.assertIsInstance(result, EvictResult)
self.assertGreaterEqual(result.num_tokens_evicted, 3)
self.assertTrue(tree.full_evictable_size() <= 3)
tree.sanity_check()
# ------- evict: mamba-only -------
def test_evict_mamba_only(self):
tree, alloc, rtp, make_req = self._build_tree()
mamba_pool = rtp.mamba_pool
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7]),
value=alloc.alloc(7),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.mamba_evictable_size(), 2)
# Evict 1 mamba state
result = tree.evict(EvictParams(num_tokens=0, mamba_num=1))
self.assertGreaterEqual(result.mamba_num_evicted, 1)
# After mamba eviction on an internal node, full tokens remain
self.assertGreaterEqual(tree.full_evictable_size(), 0)
tree.sanity_check()
# ------- evict: mamba → match stops at tombstone -------
def test_evict_mamba_breaks_match(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Evict all mamba (2 states)
tree.evict(EvictParams(num_tokens=0, mamba_num=2))
self.assertEqual(tree.mamba_evictable_size(), 0)
# Now match should return 0 because mamba validator fails
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- evict: lock_ref protection -------
def test_evict_respects_lock_ref(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12]),
value=alloc.alloc(3),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Lock the first leaf
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
locked_node = m.last_device_node
tree.inc_lock_ref(locked_node)
# Evict all full tokens — only unlocked leaf should be evicted
result = tree.evict(EvictParams(num_tokens=6))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
# [1,2,3] is still matchable because it was locked
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
self.assertEqual(len(m.device_indices), 3)
# Unlock and verify we can now evict it
tree.dec_lock_ref(locked_node)
result = tree.evict(EvictParams(num_tokens=3))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
tree.sanity_check()
# ------- evict: verify EvictResult accounting -------
def test_evict_result_accounting(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Request 0 mamba + 3 full → full evicted, mamba cascaded
result = tree.evict(EvictParams(num_tokens=3))
self.assertGreaterEqual(result.num_tokens_evicted, 3)
# Leaf eviction cascades all components; mamba also freed
self.assertGreaterEqual(result.mamba_num_evicted, 1)
tree.sanity_check()
# ------- insert: prev_prefix_len controls overlap free -------
def test_insert_prev_prefix_len(self):
tree, alloc, _, make_req = self._build_tree()
initial_avail = alloc.available_size()
# Step 1: Insert [1,2,3]
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(alloc.available_size(), initial_avail - 3)
# Step 2: Insert [1,2,3,4,5] with prev_prefix_len=0 → frees overlap [0:3]
req2 = make_req()
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=0,
)
)
self.assertEqual(result.prefix_len, 3)
# alloc 5, freed 3 overlap, stored 2 new → net -2
self.assertEqual(alloc.available_size(), initial_avail - 3 - 2)
# Step 3: Insert [1,2,3,4,5,6] with prev_prefix_len=5 → nothing freed
req3 = make_req()
avail_before = alloc.available_size()
result = tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6]),
value=alloc.alloc(6),
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=5,
)
)
self.assertEqual(result.prefix_len, 5)
# alloc 6, freed 0, stored 1 → net -6
self.assertEqual(alloc.available_size(), avail_before - 6)
tree.sanity_check()
# ------- available_and_evictable_str + pretty_print -------
def test_diagnostics(self):
tree, alloc, _, make_req = self._build_tree()
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
diag = tree.available_and_evictable_str()
self.assertIn("Available full tokens", diag)
self.assertIn("mamba", diag.lower())
diag2 = available_and_evictable_str(tree)
self.assertIn("Available full tokens", diag2)
tree.pretty_print()
tree.sanity_check()
# ===================================================================
# Test: Full + SWA + Mamba components
# ===================================================================
class TestUnifiedRadixCacheSWAMamba(unittest.TestCase):
"""UnifiedRadixCache with (Full, SWA, Mamba) components — the most complex config."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
kv_size_swa: int = 64,
max_num_reqs: int = 10,
mamba_cache_size: int = 20,
max_context_len: int = 128,
sliding_window_size: int = 4,
):
device = get_device()
with envs.SGLANG_MAMBA_SSM_DTYPE.override("bfloat16"):
shape = Mamba2StateShape.create(
tp_world_size=1,
intermediate_size=4096,
n_groups=16,
num_heads=32,
head_dim=128,
state_size=128,
conv_kernel=4,
)
mamba2_cache_params = Mamba2CacheParams(
shape=shape, layers=_mamba_layer_ids()
)
req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs,
mamba_size=mamba_cache_size,
mamba_spec_state_size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
cache_params=mamba2_cache_params,
mamba_layer_ids=_mamba_layer_ids(),
enable_mamba_extra_buffer=False,
speculative_num_draft_tokens=3,
)
kv_pool = SWAKVPool(
size=kv_size,
size_swa=kv_size_swa,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
swa_attention_layer_ids=_swa_attention_layer_ids(),
full_attention_layer_ids=_full_attention_layer_ids(),
enable_kvcache_transpose=False,
device=device,
)
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size_swa,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
device=device,
kvcache=kv_pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
sliding_window_size=sliding_window_size,
tree_components=(
ComponentType.FULL,
ComponentType.SWA,
ComponentType.MAMBA,
),
),
)
def make_req():
sp = SamplingParams(temperature=0, max_new_tokens=1)
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
sampling_params=sp,
)
req_to_token_pool.alloc([req])
return req
return tree, allocator, req_to_token_pool, make_req
# ------- basic insert + match with SWA -------
def test_insert_and_match_with_swa(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Match: SWA validator requires contiguous window >= sliding_window_size
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
# With sliding_window_size=4 and 5 tokens on single node → should match
self.assertEqual(len(m.device_indices), 5)
tree.sanity_check()
# ------- evict SWA → cascade Mamba -------
def test_evict_swa_cascades_mamba(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
# Build tree: [1,2,3] → [4,5,6,7]
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
value=alloc.alloc(3),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5, 6, 7]),
value=alloc.alloc(7),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
initial_mamba = tree.mamba_evictable_size()
# Evict SWA — on internal node, SWA eviction cascades to Mamba (priority: swa=1 > mamba=0)
result = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=3))
self.assertGreaterEqual(result.swa_num_tokens_evicted, 0)
tree.sanity_check()
# ------- evict full leaf -------
def test_evict_full_leaf_cascades_all(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12, 13, 14]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 10)
# Evict one leaf (5 full tokens) → also cascades SWA + Mamba
result = tree.evict(EvictParams(num_tokens=5))
self.assertGreaterEqual(result.num_tokens_evicted, 5)
# Leaf eviction should cascade all components
self.assertGreaterEqual(result.mamba_num_evicted, 1)
self.assertGreaterEqual(result.swa_num_tokens_evicted, 0)
tree.sanity_check()
# ------- evict with SWA lock -------
def test_swa_lock_protects_from_eviction(self):
tree, alloc, _, make_req = self._build_tree(sliding_window_size=4)
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
req2 = make_req()
tree.insert(
InsertParams(
key=RadixKey([10, 11, 12, 13, 14]),
value=alloc.alloc(5),
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
# Lock the first entry
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
lock_result = tree.inc_lock_ref(m.last_device_node)
# Try to evict all full tokens
result = tree.evict(EvictParams(num_tokens=10))
# Only the unlocked one (5 tokens) should be evictable
self.assertGreaterEqual(result.num_tokens_evicted, 5)
# Locked one is still matchable
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5])))
self.assertEqual(len(m.device_indices), 5)
# Unlock
tree.dec_lock_ref(
m.last_device_node,
DecLockRefParams(swa_uuid_for_lock=lock_result.swa_uuid_for_lock),
)
tree.sanity_check()
# ------- cache_finished_req (with insert) -------
def test_cache_finished_req_insert(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3, 4, 5]
req.output_ids = [6, 7]
kv_len = len(req.origin_input_ids) + len(req.output_ids)
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.mamba_last_track_seqlen = kv_len
req.fill_ids = req.origin_input_ids + req.output_ids
tree.cache_finished_req(req, is_insert=True)
# Verify the tokens are in the tree
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5, 6, 7])))
self.assertEqual(len(m.device_indices), 7)
tree.sanity_check()
# ------- cache_finished_req (no insert) -------
def test_cache_finished_req_no_insert(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3]
req.output_ids = []
kv_len = 3
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.fill_ids = req.origin_input_ids
avail_before = alloc.available_size()
tree.cache_finished_req(req, is_insert=False)
# KV indices should be freed back
self.assertEqual(alloc.available_size(), avail_before + kv_len)
# Nothing in tree
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- cache_unfinished_req -------
def test_cache_unfinished_req(self):
tree, alloc, rtp, make_req = self._build_tree()
req = make_req()
req.origin_input_ids = [1, 2, 3, 4, 5]
req.output_ids = []
req.fill_ids = req.origin_input_ids[:]
kv_len = len(req.fill_ids)
kv_indices = alloc.alloc(kv_len)
rtp.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
req.kv_committed_len = kv_len
req.last_node = tree.root_node
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.mamba_last_track_seqlen = kv_len
tree.cache_unfinished_req(req)
# After caching, prefix_indices should be set
self.assertGreater(len(req.prefix_indices), 0)
self.assertEqual(req.cache_protected_len, len(req.prefix_indices))
self.assertIsNotNone(req.last_node)
# Release the lock acquired by cache_unfinished_req before idle check
tree.dec_lock_ref(
req.last_node,
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
)
tree.sanity_check()
# ------- evict empty tree → no crash -------
def test_evict_empty_tree(self):
tree, alloc, _, _ = self._build_tree()
result = tree.evict(EvictParams(num_tokens=10, mamba_num=5))
self.assertEqual(result.num_tokens_evicted, 0)
self.assertEqual(result.mamba_num_evicted, 0)
tree.sanity_check()
# ------- multiple evictions until empty -------
def test_evict_until_empty(self):
tree, alloc, _, make_req = self._build_tree()
for i in range(5):
req = make_req()
tokens = list(range(i * 10, i * 10 + 5))
tree.insert(
InsertParams(
key=RadixKey(tokens),
value=alloc.alloc(5),
mamba_value=req.mamba_pool_idx.unsqueeze(0),
)
)
self.assertEqual(tree.full_evictable_size(), 25)
# Evict all
result = tree.evict(EvictParams(num_tokens=100))
self.assertGreaterEqual(result.num_tokens_evicted, 25)
self.assertEqual(tree.full_evictable_size(), 0)
self.assertEqual(tree.mamba_evictable_size(), 0)
# Verify tree is empty (no matches)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey([0, 1, 2, 3, 4])))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
# ------- cow mamba on match -------
def test_match_cow_mamba(self):
tree, alloc, rtp, make_req = self._build_tree()
mamba_pool = rtp.mamba_pool
req1 = make_req()
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
# Match with cow_mamba
req2 = make_req()
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey([1, 2, 3, 4, 5]), cow_mamba=True, req=req2)
)
self.assertEqual(len(m.device_indices), 5)
# req2 should now have its own mamba state (copied)
self.assertIsNotNone(req2.mamba_pool_idx)
# Verify the copy matches
src_value = m.last_device_node.component_data[ComponentType.MAMBA].value
self.assertTrue(
torch.all(
mamba_pool.mamba_cache.conv[0][:, req2.mamba_pool_idx]
== mamba_pool.mamba_cache.conv[0][:, src_value]
)
)
tree.sanity_check()
# ===================================================================
# Test: Helper functions
# ===================================================================
class TestUnifiedRadixCacheHelpers(unittest.TestCase):
"""Tests for internal helper functions of UnifiedRadixCache."""
@classmethod
def setUpClass(cls):
set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=_PAGE_SIZE)
)
def _build_tree(
self,
kv_size: int = 128,
max_num_reqs: int = 10,
max_context_len: int = 128,
):
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
device = get_device()
req_to_token_pool = ReqToTokenPool(
size=max_num_reqs,
max_context_len=max_context_len,
device=device,
enable_memory_saver=False,
)
kv_pool = MHATokenToKVPool(
size=kv_size,
page_size=_PAGE_SIZE,
dtype=_DTYPE,
head_num=_HEAD_NUM,
head_dim=_HEAD_DIM,
layer_num=_NUM_LAYERS,
device=device,
enable_memory_saver=False,
)
allocator = TokenToKVPoolAllocator(
size=kv_size,
dtype=_DTYPE,
device=device,
kvcache=kv_pool,
need_sort=False,
)
tree = UnifiedRadixCache(
params=CacheInitParams(
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=_PAGE_SIZE,
disable=False,
tree_components=(
ComponentType.FULL,
), # Full attention only, no mamba/swa
),
)
return tree, allocator
def test_readonly_does_not_modify_tree(self):
"""Verify readonly match does not modify tree structure (no split)."""
tree, alloc = self._build_tree()
# Insert [1, 2, 3, 4, 5]
tree.insert(
InsertParams(
key=RadixKey([1, 2, 3, 4, 5]),
value=alloc.alloc(5),
)
)
def count_nodes(node):
count = 1
for child in node.children.values():
count += count_nodes(child)
return count
node_count_before = count_nodes(tree.root_node)
self.assertEqual(node_count_before, 2) # root_node and [1, 2, 3, 4, 5]
# Regular match with partial key [1, 2] creates a split
value, best_node, best_value_len = tree._match_prefix_helper(RadixKey([1, 2]))
# Regular match with partial key [1, 2, 3, 4] creates a split
value, best_node, best_value_len = tree._match_prefix_helper(
RadixKey([1, 2, 3, 4])
)
self.assertEqual(best_value_len, 2)
self.assertEqual(best_node.key.token_ids, [3, 4])
node_count_after_regular = count_nodes(tree.root_node)
self.assertEqual(node_count_after_regular, node_count_before + 2)
# Readonly match with partial key [1, 2, 3] should NOT create a split
value, best_node, best_value_len = tree._match_prefix_helper_readonly(
RadixKey([1, 2, 3])
)
self.assertEqual(best_value_len, 1)
self.assertEqual(best_node.key.token_ids, [1, 2])
node_count_after_readonly = count_nodes(tree.root_node)
self.assertEqual(node_count_after_readonly, node_count_after_regular)
tree.sanity_check()
if __name__ == "__main__":
unittest.main()