diff --git a/benchmark/unified_memory/bench_peer_aware_eviction.py b/benchmark/unified_memory/bench_peer_aware_eviction.py new file mode 100644 index 000000000..99c635342 --- /dev/null +++ b/benchmark/unified_memory/bench_peer_aware_eviction.py @@ -0,0 +1,246 @@ +"""Reproduce and measure peer-aware eviction in the unified memory pool. + +The workload deliberately fills the shared FULL/Mamba pool with many small, +reusable prefixes. It then submits one long pressure request and probes every +prefix again. Keeping more probe prefixes cached demonstrates that allocator +capacity gained from a peer component stopped radix eviction early. + +Example: + + python benchmark/unified_memory/bench_peer_aware_eviction.py \ + --label proposed --output /tmp/proposed.json + +Concurrent fan-out from a prefix retained only by the proposed path: + + python benchmark/unified_memory/bench_peer_aware_eviction.py \ + --model-path Qwen/Qwen3.5-4B \ + --probe-output-len 32 --probe-concurrency 6 \ + --burst-target-index 14 --burst-requests 30 \ + --label proposed-concurrent --output /tmp/proposed-concurrent.json +""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Any + +import requests +from transformers import AutoTokenizer + + +def percentile(values: list[float], fraction: float) -> float: + values = sorted(values) + if not values: + return 0.0 + return values[round((len(values) - 1) * fraction)] + + +def generate(base_url: str, text: str, max_new_tokens: int) -> dict: + start = time.perf_counter() + response = requests.post( + f"{base_url}/generate", + json={ + "text": text, + "sampling_params": { + "max_new_tokens": max_new_tokens, + "temperature": 0, + "ignore_eos": True, + }, + }, + timeout=600, + ) + elapsed = time.perf_counter() - start + response.raise_for_status() + body = response.json() + meta = body["meta_info"] + prefill_finished_time = meta.get("prefill_finished_time") + forward_entry_time = meta.get("forward_entry_time") + return { + "prompt_tokens": meta["prompt_tokens"], + "cached_tokens": meta["cached_tokens"], + "e2e_latency_s": meta["e2e_latency"], + "client_latency_s": elapsed, + "prefill_latency_s": ( + prefill_finished_time - forward_entry_time + if prefill_finished_time is not None and forward_entry_time is not None + else None + ), + "completion_tokens": meta["completion_tokens"], + "num_retractions": meta["num_retractions"], + "output_ids": body["output_ids"], + } + + +def text_with_target_tokens(tokenizer, seed: str, target: int) -> str: + """Create deterministic text whose tokenized length is close to ``target``.""" + repeated = (seed + " ") * target + token_ids = tokenizer.encode(repeated, add_special_tokens=False)[:target] + return tokenizer.decode(token_ids, skip_special_tokens=True) + + +def summarize_probe(probes: list[dict], batch_wall_latency_s: float) -> dict[str, Any]: + cached = [item["cached_tokens"] for item in probes] + e2e = [item["e2e_latency_s"] for item in probes] + client = [item["client_latency_s"] for item in probes] + prefill = [ + item["prefill_latency_s"] + for item in probes + if item["prefill_latency_s"] is not None + ] + completion_tokens = sum(item["completion_tokens"] for item in probes) + return { + "cached_prefixes": sum(value > 0 for value in cached), + "cache_survival_rate": sum(value > 0 for value in cached) / len(cached), + "total_cached_tokens": sum(cached), + "mean_cached_tokens": statistics.mean(cached), + "cached_tokens": cached, + "mean_e2e_latency_s": statistics.mean(e2e), + "p95_e2e_latency_s": percentile(e2e, 0.95), + "mean_client_latency_s": statistics.mean(client), + "p95_client_latency_s": percentile(client, 0.95), + "mean_prefill_latency_s": statistics.mean(prefill) if prefill else None, + "p95_prefill_latency_s": percentile(prefill, 0.95) if prefill else None, + "batch_wall_latency_s": batch_wall_latency_s, + "request_throughput_rps": len(probes) / batch_wall_latency_s, + "output_throughput_tps": completion_tokens / batch_wall_latency_s, + "completion_tokens": completion_tokens, + "total_retractions": sum(item["num_retractions"] for item in probes), + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--base-url", default="http://127.0.0.1:30000") + parser.add_argument("--model-path", default="Qwen/Qwen3.5-0.8B") + parser.add_argument("--label", required=True) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--warm-prefixes", type=int, default=28) + parser.add_argument("--prefix-len", type=int, default=400) + parser.add_argument("--pressure-len", type=int, default=7000) + parser.add_argument("--output-len", type=int, default=1) + parser.add_argument("--probe-output-len", type=int) + parser.add_argument("--probe-concurrency", type=int, default=1) + parser.add_argument("--burst-target-index", type=int) + parser.add_argument("--burst-requests", type=int, default=30) + args = parser.parse_args() + if args.probe_concurrency < 1: + parser.error("--probe-concurrency must be at least 1") + if args.burst_requests < 1: + parser.error("--burst-requests must be at least 1") + if args.burst_target_index is not None and not ( + 0 <= args.burst_target_index < args.warm_prefixes + ): + parser.error("--burst-target-index must select a warm prefix") + probe_output_len = args.probe_output_len or args.output_len + tokenizer = AutoTokenizer.from_pretrained(args.model_path) + + prompt_pairs = [] + prompt_bases = [] + for index in range(args.warm_prefixes): + base = text_with_target_tokens( + tokenizer, + f"stable reusable unified cache prefix group {index}", + args.prefix_len - 16, + ) + prompt_bases.append(base) + prompt_pairs.append( + ( + base + f" warm suffix for group {index}", + base + f" replay suffix for group {index}", + ) + ) + pressure_text = text_with_target_tokens( + tokenizer, "distinct long allocation pressure payload", args.pressure_len + ) + + requests.post(f"{args.base_url}/flush_cache", timeout=60).raise_for_status() + + # Exclude server startup and first-request kernel initialization. + generate(args.base_url, "server kernel warmup " * 32, args.output_len) + requests.post(f"{args.base_url}/flush_cache", timeout=60).raise_for_status() + + warm = [] + for warm_text, _ in prompt_pairs: + warm.append( + generate( + args.base_url, + warm_text, + args.output_len, + ) + ) + + # A distinct long request forces the FULL side toward the Mamba frontier. + pressure = generate( + args.base_url, + pressure_text, + args.output_len, + ) + + if args.burst_target_index is None: + probe_indices = list(reversed(range(len(prompt_pairs)))) + probe_texts = [prompt_pairs[index][1] for index in probe_indices] + else: + probe_indices = [args.burst_target_index] * args.burst_requests + target_base = prompt_bases[args.burst_target_index] + probe_texts = [ + target_base + f" concurrent burst replay suffix request {ordinal}" + for ordinal in range(args.burst_requests) + ] + + def run_probe(replay_text: str) -> dict: + return generate( + args.base_url, + replay_text, + probe_output_len, + ) + + probe_start = time.perf_counter() + if args.probe_concurrency == 1: + probes = [run_probe(replay_text) for replay_text in probe_texts] + else: + with ThreadPoolExecutor(max_workers=args.probe_concurrency) as executor: + probes = list(executor.map(run_probe, probe_texts)) + probe_wall_latency_s = time.perf_counter() - probe_start + + result = { + "label": args.label, + "config": { + "warm_prefixes": args.warm_prefixes, + "prefix_len": args.prefix_len, + "pressure_len": args.pressure_len, + "output_len": args.output_len, + "probe_output_len": probe_output_len, + "probe_concurrency": args.probe_concurrency, + "burst_target_index": args.burst_target_index, + "burst_requests": ( + args.burst_requests if args.burst_target_index is not None else None + ), + "actual_warm_prompt_tokens": [item["prompt_tokens"] for item in warm], + }, + "warm": { + "total_cached_tokens": sum(item["cached_tokens"] for item in warm), + "total_retractions": sum(item["num_retractions"] for item in warm), + }, + "warm_requests": warm, + "pressure": pressure, + "probe": summarize_probe(probes, probe_wall_latency_s), + "probe_requests": probes, + "probe_indices": probe_indices, + "output_ids": { + "warm": [item["output_ids"] for item in warm], + "pressure": pressure["output_ids"], + "probe": [item["output_ids"] for item in probes], + }, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(result, indent=2) + "\n") + print(json.dumps(result["probe"], indent=2)) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 667294da4..e49da8ee4 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -427,7 +427,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): required = ceil_align(swa_tail_len, page_size) available = self.token_to_kv_pool_allocator.swa_available_size() if available < required: - self.tree_cache.evict(EvictParams(swa_num_tokens=required - available)) + self.tree_cache.evict_for_alloc( + EvictParams(swa_num_tokens=required - available) + ) available = self.token_to_kv_pool_allocator.swa_available_size() if available < required: @@ -1789,7 +1791,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): and self._radix_full_available() < required_alloc_tokens ): num_to_evict = required_alloc_tokens - self._radix_full_available() - result = self.tree_cache.evict(EvictParams(num_tokens=num_to_evict)) + result = self.tree_cache.evict_for_alloc( + EvictParams(num_tokens=num_to_evict) + ) if self._radix_full_available() < required_alloc_tokens: logger.warning( f"Eviction insufficient: needed {required_alloc_tokens} tokens, " diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py index 8fe3b4827..a5abd9955 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/auxiliary_state.py @@ -10,7 +10,7 @@ storing model-agnostic native cache snapshots. from __future__ import annotations from dataclasses import dataclass -from typing import Any, Iterable, Optional +from typing import Any, Iterable, Iterator, Optional import mlx.core as mx import torch @@ -95,6 +95,7 @@ class MlxAuxiliaryStatePool: self.mamba_cache = None self.mem_usage = 0 self._snapshots: dict[int, dict[int, _CacheSnapshot]] = {} + self._alloc_iter: Optional[Iterator[torch.Tensor]] = None self.clear() def _tensor(self, indices: Any) -> torch.Tensor: @@ -108,7 +109,31 @@ class MlxAuxiliaryStatePool: def available_size(self) -> int: return int(self.free_slots.numel()) + def schedulable_available_size(self) -> int: + return self.available_size() + + def alloc_group_begin(self, num_reqs: int) -> None: + self._alloc_iter = None + if num_reqs > 0: + slots = self._do_alloc(num_reqs) + if slots is not None: + self._alloc_iter = iter(slots.split(1)) + + def alloc_group_end(self) -> None: + if self._alloc_iter is not None: + remaining = list(self._alloc_iter) + if remaining: + self.free(torch.cat(remaining)) + self._alloc_iter = None + def alloc(self, need_size: int) -> Optional[torch.Tensor]: + if self._alloc_iter is not None and need_size == 1: + slot = next(self._alloc_iter, None) + if slot is not None: + return slot + return self._do_alloc(need_size) + + def _do_alloc(self, need_size: int) -> Optional[torch.Tensor]: if need_size > self.available_size(): return None slots = self.free_slots[:need_size].clone() @@ -128,6 +153,7 @@ class MlxAuxiliaryStatePool: self.free_slots = torch.cat([self.free_slots, indices]) def clear(self) -> None: + self._alloc_iter = None self.free_slots = torch.arange( 1, self.size + 1, dtype=torch.int64, device=self.device ) @@ -227,6 +253,7 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool): size=auxiliary_state_size, device=device, ) + self.mamba_allocator = self.mamba_pool # The unified radix base MAMBA component still reads ``mamba_pool``. # Keep the MLX-owned name beside it so local code can avoid model- # specific terminology. @@ -352,7 +379,7 @@ class MlxAuxiliaryStateComponent(MambaComponent): source_value ) if forked_value is None: - self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1)) forked_value = ( self.cache.req_to_token_pool.auxiliary_state_pool.fork_from( source_value diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index fe20022aa..259be8837 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -257,7 +257,9 @@ def alloc_req_slots( if mamba_available_size < mamba_state_needed: if tree_cache is not None and tree_cache.supports_mamba(): mamba_num = max(0, mamba_state_needed - mamba_available_size) - tree_cache.evict(EvictParams(num_tokens=0, mamba_num=mamba_num)) + tree_cache.evict_for_alloc( + EvictParams(num_tokens=0, mamba_num=mamba_num) + ) req_pool_indices = req_to_token_pool.alloc(reqs) if req_pool_indices is None: raise RuntimeError( diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index abc840cda..15aee4c93 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -324,6 +324,16 @@ class BasePrefixCache(ABC, PrefixCacheTrait): def evict(self, params: EvictParams) -> EvictResult: pass + def evict_for_alloc(self, params: EvictParams) -> EvictResult: + """Evict cache entries to cover allocator shortfalls. + + The default implementation preserves the component-count semantics of + :meth:`evict`. Multi-component caches backed by shared memory can + override this entry point to stop once collateral frees make the + requested allocation feasible. + """ + return self.evict(params) + @abstractmethod def inc_lock_ref(self, node: Any) -> IncLockRefResult: pass diff --git a/python/sglang/srt/mem_cache/buffer_mode/pipeline.py b/python/sglang/srt/mem_cache/buffer_mode/pipeline.py index f67efdfb5..a2facd33c 100644 --- a/python/sglang/srt/mem_cache/buffer_mode/pipeline.py +++ b/python/sglang/srt/mem_cache/buffer_mode/pipeline.py @@ -832,8 +832,12 @@ class BufferModePipeline: avail = cache.token_to_kv_pool_allocator.available_size() if avail < f.num_tokens: needed = f.num_tokens - avail - evicted = cache.evict(EvictParams(num_tokens=needed)) - if evicted.num_tokens_evicted < needed: + cache.evict_for_alloc(EvictParams(num_tokens=needed)) + if cache.supports_swa(): + avail = cache.token_to_kv_pool_allocator.full_available_size() + else: + avail = cache.token_to_kv_pool_allocator.available_size() + if avail < f.num_tokens: # Genuinely no room (locked pages): recompute. return _drop() diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 8b92cfe48..524950b3f 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -130,14 +130,16 @@ def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int): if full_available_size < num_tokens or swa_available_size < num_tokens: full_num_tokens = max(0, num_tokens - full_available_size) swa_num_tokens = max(0, num_tokens - swa_available_size) - tree_cache.evict( + tree_cache.evict_for_alloc( EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens) ) else: # Standard allocator: evict only the shortfall (mirrors the SWA arm) available_size = allocator.available_size() if available_size < num_tokens: - tree_cache.evict(EvictParams(num_tokens=num_tokens - available_size)) + tree_cache.evict_for_alloc( + EvictParams(num_tokens=num_tokens - available_size) + ) def retraction_backup( diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index ad52d2bef..b1900ea65 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -48,6 +48,26 @@ def _get_allocator_type(server_args: ServerArgs) -> str: return get_allocator_type(server_args) +def _evict_swa_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None: + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + + available_size = cache.token_to_kv_pool_allocator.swa_available_size() + shortfall = max(0, required_size - available_size) + if shortfall > 0: + cache.evict_for_alloc(EvictParams(swa_num_tokens=shortfall)) + + +def _evict_mamba_for_device_alloc(cache: UnifiedRadixCache, required_size: int) -> None: + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + + available_size = ( + cache.req_to_token_pool.mamba_allocator.schedulable_available_size() + ) + shortfall = max(0, required_size - available_size) + if shortfall > 0: + cache.evict_for_alloc(EvictParams(mamba_num=shortfall)) + + def _make_layer_mapper( layer_mapping: dict[int, int], transfer_layer_num: int, @@ -1210,8 +1230,6 @@ class _DeepSeekV4Strategy(StackStrategy): model_name=None, enable_storage_metrics=False, ): - from sglang.srt.mem_cache.base_prefix_cache import EvictParams - host_pool_group, cache_controller = build_deepseek_v4_hicache_stack( params=params, server_args=server_args, @@ -1219,7 +1237,7 @@ class _DeepSeekV4Strategy(StackStrategy): load_cache_event=load_cache_event, storage_backend=storage_backend, host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), - device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)), + device_swa_evict_fn=lambda n: _evict_swa_for_device_alloc(cache, n), prefetch_threshold=prefetch_threshold, model_name=model_name, storage_backend_extra_config=storage_backend_extra_config, @@ -1286,8 +1304,6 @@ class _MambaStrategy(StackStrategy): model_name=None, enable_storage_metrics=False, ): - from sglang.srt.mem_cache.base_prefix_cache import EvictParams - full_layer_mapping = dict(kvcache.full_attention_layer_id_mapping) mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) host_pool_group, cache_controller = build_hybrid_mamba_stack( @@ -1301,7 +1317,7 @@ class _MambaStrategy(StackStrategy): storage_backend=storage_backend, use_mla=kvcache.use_mla, host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA), - device_mamba_evict_fn=lambda n: cache.evict(EvictParams(mamba_num=n)), + device_mamba_evict_fn=lambda n: _evict_mamba_for_device_alloc(cache, n), prefetch_threshold=prefetch_threshold, model_name=model_name, storage_backend_extra_config=storage_backend_extra_config, @@ -1355,8 +1371,6 @@ class _SwaStrategy(StackStrategy): model_name=None, enable_storage_metrics=False, ): - from sglang.srt.mem_cache.base_prefix_cache import EvictParams - full_layer_mapping, swa_layer_mapping = _swa_layer_mappings(kvcache) host_pool_group, cache_controller = build_hybrid_swa_stack( params=params, @@ -1369,7 +1383,7 @@ class _SwaStrategy(StackStrategy): storage_backend=storage_backend, use_mla=False, host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), - device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)), + device_swa_evict_fn=lambda n: _evict_swa_for_device_alloc(cache, n), prefetch_threshold=prefetch_threshold, model_name=model_name, storage_backend_extra_config=storage_backend_extra_config, @@ -1417,8 +1431,6 @@ class _MambaSwaStrategy(StackStrategy): model_name=None, enable_storage_metrics=False, ): - from sglang.srt.mem_cache.base_prefix_cache import EvictParams - full_layer_mapping, swa_layer_mapping = _swa_layer_mappings(kvcache) mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) host_pool_group, cache_controller = build_hybrid_mamba_swa_stack( @@ -1438,9 +1450,9 @@ class _MambaSwaStrategy(StackStrategy): pp_group=params.pp_cache_group, storage_backend=storage_backend, host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), - device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)), + device_swa_evict_fn=lambda n: _evict_swa_for_device_alloc(cache, n), host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA), - device_mamba_evict_fn=lambda n: cache.evict(EvictParams(mamba_num=n)), + device_mamba_evict_fn=lambda n: _evict_mamba_for_device_alloc(cache, n), prefetch_threshold=prefetch_threshold, model_name=model_name, storage_backend_extra_config=storage_backend_extra_config, diff --git a/python/sglang/srt/mem_cache/pool_host/group.py b/python/sglang/srt/mem_cache/pool_host/group.py index 498b60696..55f01b139 100644 --- a/python/sglang/srt/mem_cache/pool_host/group.py +++ b/python/sglang/srt/mem_cache/pool_host/group.py @@ -16,6 +16,8 @@ class PoolEntry: device_pool: Any layer_mapper: Callable[[int], int | None] is_primary_index_anchor: bool = False + # Reclaim callbacks receive the absolute allocation size n. The host + # callback evicts n slots; the device callback makes alloc(n) feasible. host_evict_fn: Callable[[int], Any] | None = None device_evict_fn: Callable[[int], Any] | None = None device_alloc_fn: Callable[[int], Any] | None = None diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py index 3fbd7ad2b..319c8c936 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py @@ -203,7 +203,7 @@ class MambaComponent(TreeComponent): # stops at this request's window boundary instead of walking to # root and over-decrementing locks held by other requests. lock_result = self.cache.inc_lock_ref(result.best_match_node) - self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1)) dst_index = self.cache.req_to_token_pool.mamba_allocator.alloc(1) self.cache.dec_lock_ref( result.best_match_node, lock_result.to_dec_params() @@ -374,10 +374,14 @@ class MambaComponent(TreeComponent): device_frees: dict[ComponentType, list[torch.Tensor]], host_frees: dict[ComponentType, list[torch.Tensor]], ) -> Optional[NodeId]: - """Return the next device-leaf node for the driver to evict, or None. - Internal nodes are tombstoned inline (no IO). If the previous node's - eviction removed the cursor, the walk resumes from the partition - sentinel with session refs on, else it restarts at the LRU tail.""" + """Advance one device-eviction step and return a leaf, if selected. + + An internal tombstone is one complete step so the caller can apply its + pending frees and recheck allocator capacity before the next mutation. + If the previous node's eviction removed the cursor, the walk resumes + from the partition sentinel with session refs on, else it restarts at + the LRU tail. + """ ct = self.component_type lru = self.tree_core.lru_lists[ct] enabled = self.tree_core.enable_session_radix_cache @@ -387,34 +391,36 @@ class MambaComponent(TreeComponent): self._evict_device_cursor = ( lru.cursor_next() if enabled else lru.get_lru_no_lock() ) - while ( - tracker[ct] < self._evict_device_request_cnt - and self._evict_device_cursor is not None - and lru.in_list(self._evict_device_cursor) + if ( + tracker[ct] >= self._evict_device_request_cnt + or self._evict_device_cursor is None + or not lru.in_list(self._evict_device_cursor) ): - x = self._evict_device_cursor - assert x.component_data[ct].value is not None - if x in self.tree_core.evictable_device_leaves and ( - not enabled or self._can_evict_leaf_atomically(x) - ): - self._evict_device_cursor = ( - lru.cursor_next() if enabled else lru.get_prev_no_lock(x) - ) - return x.id - if not enabled: - x_next = lru.get_prev_no_lock(x) - self.tree_core._evict_component_and_detach_lru( - x, - self, - target=EvictLayer.DEVICE, - tracker=tracker, - device_frees=device_frees, - host_frees=host_frees, + return None + + x = self._evict_device_cursor + assert x.component_data[ct].value is not None + if x in self.tree_core.evictable_device_leaves and ( + not enabled or self._can_evict_leaf_atomically(x) + ): + self._evict_device_cursor = ( + lru.cursor_next() if enabled else lru.get_prev_no_lock(x) ) - self.tree_core._cascade_evict( - x, self, tracker, device_frees=device_frees, host_frees=host_frees - ) - self._evict_device_cursor = lru.cursor_next() if enabled else x_next + return x.id + if not enabled: + x_next = lru.get_prev_no_lock(x) + self.tree_core._evict_component_and_detach_lru( + x, + self, + target=EvictLayer.DEVICE, + tracker=tracker, + device_frees=device_frees, + host_frees=host_frees, + ) + self.tree_core._cascade_evict( + x, self, tracker, device_frees=device_frees, host_frees=host_frees + ) + self._evict_device_cursor = lru.cursor_next() if enabled else x_next return None def _evict_device_end(self) -> None: @@ -487,7 +493,7 @@ class MambaComponent(TreeComponent): """Allocate one mamba pool slot, evicting if necessary.""" slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1) if slot is None: - self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1)) slot = self.cache.req_to_token_pool.mamba_allocator.alloc(1) assert slot is not None, "Can not alloc mamba cache" return slot @@ -660,7 +666,7 @@ class MambaComponent(TreeComponent): return PrepareLoadBackResult() dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1) if dst is None: - self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + self.cache.evict_for_alloc(EvictParams(num_tokens=0, mamba_num=1)) dst = self.cache.req_to_token_pool.mamba_allocator.alloc(1) assert dst is not None, "Cannot alloc mamba for load_back" req.mamba_pool_idx = dst[0] diff --git a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py index a252541d2..49843b525 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py @@ -535,10 +535,14 @@ class SWAComponent(TreeComponent): device_frees: dict[ComponentType, list[torch.Tensor]], host_frees: dict[ComponentType, list[torch.Tensor]], ) -> Optional[NodeId]: - """Return the next device-leaf node for the driver to evict, or None. - Internal nodes are tombstoned inline (no IO). If the previous node's - eviction removed the cursor, the walk resumes from the partition - sentinel with session refs on, else it restarts at the LRU tail.""" + """Advance one device-eviction step and return a leaf, if selected. + + An internal tombstone is one complete step so the caller can apply its + pending frees and recheck allocator capacity before the next mutation. + If the previous node's eviction removed the cursor, the walk resumes + from the partition sentinel with session refs on, else it restarts at + the LRU tail. + """ ct = self.component_type lru = self.tree_core.lru_lists[ct] enabled = self.tree_core.enable_session_radix_cache @@ -548,34 +552,36 @@ class SWAComponent(TreeComponent): self._evict_device_cursor = ( lru.cursor_next() if enabled else lru.get_lru_no_lock() ) - while ( - tracker[ct] < self._evict_device_request_cnt - and self._evict_device_cursor is not None - and lru.in_list(self._evict_device_cursor) + if ( + tracker[ct] >= self._evict_device_request_cnt + or self._evict_device_cursor is None + or not lru.in_list(self._evict_device_cursor) ): - x = self._evict_device_cursor - assert x.component_data[ct].value is not None - if x in self.tree_core.evictable_device_leaves and ( - not enabled or self._can_evict_leaf_atomically(x) - ): - self._evict_device_cursor = ( - lru.cursor_next() if enabled else lru.get_prev_no_lock(x) - ) - return x.id - if not enabled: - x_next = lru.get_prev_no_lock(x) - self.tree_core._evict_component_and_detach_lru( - x, - self, - target=EvictLayer.DEVICE, - tracker=tracker, - device_frees=device_frees, - host_frees=host_frees, + return None + + x = self._evict_device_cursor + assert x.component_data[ct].value is not None + if x in self.tree_core.evictable_device_leaves and ( + not enabled or self._can_evict_leaf_atomically(x) + ): + self._evict_device_cursor = ( + lru.cursor_next() if enabled else lru.get_prev_no_lock(x) ) - self.tree_core._cascade_evict( - x, self, tracker, device_frees=device_frees, host_frees=host_frees - ) - self._evict_device_cursor = lru.cursor_next() if enabled else x_next + return x.id + if not enabled: + x_next = lru.get_prev_no_lock(x) + self.tree_core._evict_component_and_detach_lru( + x, + self, + target=EvictLayer.DEVICE, + tracker=tracker, + device_frees=device_frees, + host_frees=host_frees, + ) + self.tree_core._cascade_evict( + x, self, tracker, device_frees=device_frees, host_frees=host_frees + ) + self._evict_device_cursor = lru.cursor_next() if enabled else x_next return None def _evict_device_end(self) -> None: diff --git a/python/sglang/srt/mem_cache/unified_cache/components/tree_component.py b/python/sglang/srt/mem_cache/unified_cache/components/tree_component.py index 44de7926d..53f474926 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/tree_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/tree_component.py @@ -507,8 +507,11 @@ class TreeComponent(ABC): device_frees: dict[ComponentType, list[torch.Tensor]], host_frees: dict[ComponentType, list[torch.Tensor]], ) -> Optional[NodeId]: - """Return the next device-leaf node for the driver to evict, or None. - Internal nodes are tombstoned inline (no IO).""" + """Advance one eviction step and return a device leaf, if selected. + + Implementations must return after one allocator-relevant internal + mutation so the caller can drain pending frees before continuing. + """ assert ( self.is_evict_device_ongoing ), f"{self.component_type} device eviction not started" @@ -534,7 +537,7 @@ class TreeComponent(ABC): device_frees: dict[ComponentType, list[torch.Tensor]], host_frees: dict[ComponentType, list[torch.Tensor]], ) -> Optional[NodeId]: - """Advance the walk; return the next device leaf or None.""" + """Advance the walk by at most one allocator-relevant mutation.""" ... @abstractmethod diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py index faa14da3e..a50d4fd27 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py @@ -1224,7 +1224,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): def evict_device_next_node( self, component_type: ComponentType, tracker: dict[ComponentType, int] ) -> EvictDeviceNextNodeResult: - """Return the next device leaf to evict for a component, or None when done.""" + """Advance one component eviction step and report whether it progressed.""" result = EvictDeviceNextNodeResult() # The walk reads running totals for its doneness check; the result # carries only this step's delta. @@ -1236,6 +1236,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): delta = n - tracker.get(ct, 0) if delta: result.tracker[ct] = delta + result.made_progress = result.node_id is not None or bool(result.tracker) return result def evict_device_end(self, component_type: ComponentType) -> None: diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py index 03c0c65fb..53b2f420f 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py @@ -41,7 +41,15 @@ class BaseEvictionResult(msgspec.Struct): class EvictDeviceNextNodeResult(BaseEvictionResult): + """One device-walk step. + + ``node_id`` selects a leaf for the Controller to evict. ``made_progress`` + also covers an internal tombstone that returned no leaf, distinguishing it + from true walk exhaustion. + """ + node_id: Optional[NodeId] = None + made_progress: bool = False class EvictDeviceLeafResult(BaseEvictionResult): @@ -230,8 +238,11 @@ class UnifiedTreeCoreInterface(ABC): def evict_device_next_node( self, component_type: ComponentType, tracker: dict[ComponentType, int] ) -> EvictDeviceNextNodeResult: - """The next evictable node (None node_id when the walk is exhausted); - tracker is the caller's running totals, read for the doneness check.""" + """Advance one eviction step. + + A missing ``node_id`` is exhausted only when ``made_progress`` is also + false. ``tracker`` is the caller's running totals, read for doneness. + """ ... @abstractmethod diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 97f64df1a..4f9bd1240 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -531,18 +531,68 @@ class UnifiedRadixCache(BasePrefixCache): self._apply_cache_actions(self.tree_core.end_insert()) def evict(self, params: EvictParams) -> EvictResult: + return self._evict(params) + + def evict_for_alloc(self, params: EvictParams) -> EvictResult: + """Evict until the requested component allocations become feasible. + + ``params`` contains allocator shortfalls, not absolute eviction quotas. + A component eviction can cascade to its peers; with a shared memory pool, + those collateral frees can satisfy the original allocation before the + triggering component's requested count is reached. + """ if self.disable: return EvictResult() - start_time = time.perf_counter() - tracker = {ct: 0 for ct in self.tree_components} - request_by_type = { + request_by_type = self._evict_request_by_type(params) + available_size_targets = { + ct: self._component_available_size(ct) + request_cnt + for ct, request_cnt in request_by_type.items() + if request_cnt > 0 + } + return self._evict(params, available_size_targets) + + @staticmethod + def _evict_request_by_type(params: EvictParams) -> dict[ComponentType, int]: + return { ComponentType.FULL: params.num_tokens, ComponentType.SWA: params.swa_num_tokens, ComponentType.MAMBA: params.mamba_num, ComponentType.C128: 0, } - self._evict_components(request_by_type, tracker) + + def _component_available_size(self, component_type: ComponentType) -> int: + """Return capacity usable by the component's next allocation. + + Shared allocators expose schedulable capacity, which includes peer holes + that an urgent allocator flush can reclaim without further eviction. + """ + if component_type == ComponentType.FULL: + if self.supports_swa(): + return self.token_to_kv_pool_allocator.full_available_size() + return self.token_to_kv_pool_allocator.available_size() + if component_type == ComponentType.SWA: + return self.token_to_kv_pool_allocator.swa_available_size() + if component_type == ComponentType.MAMBA: + return self.req_to_token_pool.mamba_allocator.schedulable_available_size() + raise ValueError(f"Unsupported cache component: {component_type}") + + def _evict( + self, + params: EvictParams, + available_size_targets: Optional[dict[ComponentType, int]] = None, + ) -> EvictResult: + if self.disable: + return EvictResult() + start_time = time.perf_counter() + tracker = {ct: 0 for ct in self.tree_components} + + request_by_type = self._evict_request_by_type(params) + self._evict_components( + request_by_type, + tracker, + available_size_targets=available_size_targets, + ) if ( self.cache_controller is not None @@ -581,12 +631,12 @@ class UnifiedRadixCache(BasePrefixCache): def _evict_device_next_node( self, component_type: ComponentType, tracker: dict[ComponentType, int] - ) -> Optional[NodeId]: + ) -> tuple[Optional[NodeId], bool]: """Advance the eviction walk one node, consuming its step result.""" result = self.tree_core.evict_device_next_node(component_type, tracker) self._free_values(result.device_frees, result.host_frees) self._accumulate_tracker(tracker, result.tracker) - return result.node_id + return result.node_id, result.made_progress def _evict_device_leaf( self, node_id: NodeId, tracker: dict[ComponentType, int] @@ -617,20 +667,39 @@ class UnifiedRadixCache(BasePrefixCache): self, request_by_type: dict[ComponentType, int], tracker: dict[ComponentType, int], + available_size_targets: Optional[dict[ComponentType, int]] = None, ) -> None: # Buffer mode: eviction always wins over queued backup intents — a # destroyed victim's intent is stale-swept and the content rewrites # after its recompute. + + def target_reached(component_type: ComponentType) -> bool: + if available_size_targets is None: + return False + target = available_size_targets.get(component_type) + # Do not compact on every eviction step. Shared allocators include + # drainable peer holes here and flush the peer once in alloc(). + return ( + target is not None + and self._component_available_size(component_type) >= target + ) + for ct in self.tree_components: request_cnt = request_by_type[ct] - # Skip eviction walk if request is already met - if tracker[ct] >= request_cnt: + # A preceding component may have cascade-evicted this component or, + # on a shared pool, released enough bytes to satisfy its allocation. + if tracker[ct] >= request_cnt or target_reached(ct): continue self.tree_core.evict_device_start(ct, request_cnt) try: - while ( - node_id := self._evict_device_next_node(ct, tracker) - ) is not None: + while not target_reached(ct): + node_id, made_progress = self._evict_device_next_node(ct, tracker) + if node_id is None: + if made_progress: + # Internal tombstone frees are now allocator-visible; + # recheck the allocation target before walking again. + continue + break backup_kv = self._evict_device_leaf(node_id, tracker) if backup_kv is not None: # Deferred demote: run the D->H backup, demote only on success. @@ -1395,14 +1464,11 @@ class UnifiedRadixCache(BasePrefixCache): self.dec_host_lock_ref(node_id, host_anchor_params) return False - if self.supports_swa(): - avail = self.token_to_kv_pool_allocator.full_available_size() - else: - avail = self.token_to_kv_pool_allocator.available_size() + avail = self._component_available_size(ComponentType.FULL) if avail < kv_tokens: needed = kv_tokens - avail - result = self.evict(EvictParams(num_tokens=needed)) - if result.num_tokens_evicted < needed: + self.evict_for_alloc(EvictParams(num_tokens=needed)) + if self._component_available_size(ComponentType.FULL) < kv_tokens: self.dec_lock_ref(node_id, ancestor_lock_params) self.dec_host_lock_ref(node_id, host_anchor_params) return False diff --git a/python/sglang/srt/session/streaming_session.py b/python/sglang/srt/session/streaming_session.py index bdbf38d75..9489b1e26 100644 --- a/python/sglang/srt/session/streaming_session.py +++ b/python/sglang/srt/session/streaming_session.py @@ -417,6 +417,9 @@ class StreamingSession(BasePrefixCache): def evict(self, params: EvictParams) -> EvictResult: return self.inner.evict(params) + def evict_for_alloc(self, params: EvictParams) -> EvictResult: + return self.inner.evict_for_alloc(params) + def inc_lock_ref(self, node: Any) -> IncLockRefResult: result = self.try_inc_lock_ref(node) if result is not None: diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index fb9d1e5ab..14dc72d4e 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -877,6 +877,17 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): self.assertEqual(forked.tolist(), [3]) self.assertEqual(restored[0].state[0].tolist(), [1.0]) self.assertEqual(pool.available_size(), 3) + self.assertEqual(pool.schedulable_available_size(), 3) + + def test_auxiliary_state_pool_returns_unused_group_slots(self): + pool = MlxAuxiliaryStatePool(size=4, device="cpu") + + pool.alloc_group_begin(3) + allocated = pool.alloc(1) + pool.alloc_group_end() + + self.assertEqual(allocated.tolist(), [1]) + self.assertEqual(pool.available_size(), 3) def test_auxiliary_state_pool_restores_instance_meta_state(self): pool = MlxAuxiliaryStatePool(size=2, device="cpu") @@ -916,6 +927,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): self.assertIsNotNone(auxiliary_state_idx) self.assertIsNone(req.req_pool_idx) self.assertIsNotNone(req.mamba_pool_idx) + self.assertIs(pool.mamba_allocator, pool.mamba_pool) self.assertEqual(pool.auxiliary_state_pool.available_size(), 3) pool.free_auxiliary_state_cache(req) self.assertIsNone(req.mamba_pool_idx) diff --git a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py index e9f2116bf..4216f1998 100644 --- a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py +++ b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py @@ -144,7 +144,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase): error = queue._reclaim_swa_tail_capacity(129, "req-1") self.assertIsNone(error) - params = queue.tree_cache.evict.call_args.args[0] + params = queue.tree_cache.evict_for_alloc.call_args.args[0] self.assertEqual(params.num_tokens, 0) self.assertEqual(params.swa_num_tokens, 128) diff --git a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py index 3849e7573..185abd5f8 100644 --- a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py @@ -2,9 +2,12 @@ import unittest from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch +from sglang.srt.mem_cache.base_prefix_cache import EvictParams from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( + _evict_mamba_for_device_alloc, + _evict_swa_for_device_alloc, _split_hicache_size, build_full_draft_pools, ) @@ -23,6 +26,40 @@ class _Pool: return self._kv_bytes +class TestDeviceAllocEviction(CustomTestCase): + def test_swa_evicts_only_allocation_shortfall(self): + cache = MagicMock() + cache.token_to_kv_pool_allocator.swa_available_size.return_value = 8 + + _evict_swa_for_device_alloc(cache, required_size=10) + + cache.evict_for_alloc.assert_called_once_with(EvictParams(swa_num_tokens=2)) + cache.evict.assert_not_called() + + def test_mamba_evicts_only_allocation_shortfall(self): + cache = MagicMock() + allocator = cache.req_to_token_pool.mamba_allocator + allocator.schedulable_available_size.return_value = 8 + + _evict_mamba_for_device_alloc(cache, required_size=10) + + cache.evict_for_alloc.assert_called_once_with(EvictParams(mamba_num=2)) + cache.evict.assert_not_called() + + def test_sufficient_capacity_skips_eviction(self): + cache = MagicMock() + cache.token_to_kv_pool_allocator.swa_available_size.return_value = 10 + cache.req_to_token_pool.mamba_allocator.schedulable_available_size.return_value = ( + 10 + ) + + _evict_swa_for_device_alloc(cache, required_size=10) + _evict_mamba_for_device_alloc(cache, required_size=10) + + cache.evict_for_alloc.assert_not_called() + cache.evict.assert_not_called() + + class TestSplitHicacheSize(CustomTestCase): def test_splits_total_budget_by_device_bytes(self): # scalar and (k, v) tuple return shapes both supported diff --git a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py index 67b6e47c8..1fd938b2a 100644 --- a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py +++ b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py @@ -55,10 +55,10 @@ class _RatioCache: self.component_evictable_size_ = {ComponentType.MAMBA: 0} self.component_protected_size_ = {ComponentType.MAMBA: 0} self.prefix_nodes = [] + self.alloc_evict_params = [] - def evict(self, params: EvictParams): - # Reclaim up to mamba_num evictable (unlocked) prefix snapshots, mirroring - # what the real tree eviction can hand back under mamba pressure. + def evict_for_alloc(self, params: EvictParams): + self.alloc_evict_params.append(params) need = params.mamba_num for node in list(self.prefix_nodes): if need <= 0: @@ -130,7 +130,9 @@ class TestMambaRatioEnvGate(unittest.TestCase): return KVCacheConfigurator._calculate_mamba_ratio(fake) def test_flag_off_restores_original_ratios(self): - r = lambda **kw: self._ratio(skip=False, **kw) + def r(**kwargs): + return self._ratio(skip=False, **kwargs) + self.assertEqual( r(extra_buffer=False, lazy=False, disable_overlap=True), 3 ) # no_buffer @@ -142,7 +144,9 @@ class TestMambaRatioEnvGate(unittest.TestCase): ) # overlap def test_flag_on_drops_base_but_keeps_no_buffer(self): - r = lambda **kw: self._ratio(skip=True, **kw) + def r(**kwargs): + return self._ratio(skip=True, **kwargs) + self.assertEqual( r(extra_buffer=False, lazy=False, disable_overlap=True), 3 ) # no_buffer @@ -212,9 +216,12 @@ class TestDecSwaLockSkip(unittest.TestCase): class TestMambaDonatedAllocRatio(unittest.TestCase): def test_prefill_peak_ratio2_exhausts_pool(self): # pool = 2N, all N prefixes admission-locked: no evictable victim. - component, _, _ = _build_peak(pool_size=2 * N, lock_prefixes=True) + component, cache, _ = _build_peak(pool_size=2 * N, lock_prefixes=True) with self.assertRaisesRegex(AssertionError, "Can not alloc mamba cache"): component._alloc_mamba_slot() + self.assertEqual( + cache.alloc_evict_params, [EvictParams(num_tokens=0, mamba_num=1)] + ) def test_prefill_peak_ratio3_has_headroom(self): # pool = 3N: N free slots remain after own + locked prefix. @@ -230,6 +237,9 @@ class TestMambaDonatedAllocRatio(unittest.TestCase): slot = component._alloc_mamba_slot() self.assertIsNotNone(slot) self.assertEqual(len(cache.prefix_nodes), N - 1) + self.assertEqual( + cache.alloc_evict_params, [EvictParams(num_tokens=0, mamba_num=1)] + ) class TestPPMambaPoolSizing(unittest.TestCase): diff --git a/test/registered/unit/mem_cache/test_unified_radix_allocation_eviction.py b/test/registered/unit/mem_cache/test_unified_radix_allocation_eviction.py new file mode 100644 index 000000000..669059511 --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_radix_allocation_eviction.py @@ -0,0 +1,132 @@ +"""CPU-only tests for allocation-aware UnifiedRadixCache eviction.""" + +import unittest +from unittest.mock import MagicMock + +from sglang.srt.mem_cache.base_prefix_cache import EvictParams +from sglang.srt.mem_cache.common import evict_from_tree_cache +from sglang.srt.mem_cache.unified_cache.components import ComponentType +from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestUnifiedRadixAllocationEviction(CustomTestCase): + @staticmethod + def _build_cache(*, collateral_capacity_gain: int): + cache = object.__new__(UnifiedRadixCache) + cache.disable = False + cache.tree_components = (ComponentType.FULL, ComponentType.MAMBA) + cache.is_swa_enabled = False + cache.cache_controller = None + cache.metrics_collector = None + cache.tree_core = MagicMock() + + capacity = {"available": 30} + allocator = MagicMock() + allocator.available_size.side_effect = lambda: capacity["available"] + cache.token_to_kv_pool_allocator = allocator + cache.req_to_token_pool = MagicMock() + + leaf_count = {"value": 0} + + def next_node(component_type, tracker): + if tracker[component_type] >= 70: + return None, False + return leaf_count["value"] + 1, True + + def evict_leaf(_node_id, tracker): + leaf_count["value"] += 1 + tracker[ComponentType.FULL] += 20 + tracker[ComponentType.MAMBA] += 1 + capacity["available"] += ( + collateral_capacity_gain if leaf_count["value"] == 1 else 20 + ) + return None + + cache._evict_device_next_node = MagicMock(side_effect=next_node) + cache._evict_device_leaf = MagicMock(side_effect=evict_leaf) + return cache, capacity, leaf_count + + def test_allocation_eviction_stops_when_shared_capacity_is_sufficient(self): + cache, capacity, leaf_count = self._build_cache(collateral_capacity_gain=70) + + result = cache.evict_for_alloc(EvictParams(num_tokens=70)) + + self.assertEqual(capacity["available"], 100) + self.assertEqual(leaf_count["value"], 1) + self.assertEqual(result.num_tokens_evicted, 20) + self.assertEqual(result.mamba_num_evicted, 1) + + def test_explicit_evict_preserves_component_count_semantics(self): + cache, _, leaf_count = self._build_cache(collateral_capacity_gain=70) + + result = cache.evict(EvictParams(num_tokens=70)) + + self.assertEqual(leaf_count["value"], 4) + self.assertEqual(result.num_tokens_evicted, 80) + self.assertEqual(result.mamba_num_evicted, 4) + + def test_c128_component_keeps_zero_quota(self): + cache, _, _ = self._build_cache(collateral_capacity_gain=70) + cache.tree_components = (ComponentType.FULL, ComponentType.C128) + cache._evict_device_next_node.side_effect = None + cache._evict_device_next_node.return_value = (None, False) + + result = cache.evict(EvictParams(num_tokens=1)) + + self.assertEqual(result.num_tokens_evicted, 0) + cache.tree_core.evict_device_start.assert_called_once_with( + ComponentType.FULL, 1 + ) + + def test_mamba_allocation_counts_collateral_full_capacity(self): + cache = object.__new__(UnifiedRadixCache) + cache.disable = False + cache.tree_components = (ComponentType.FULL, ComponentType.MAMBA) + cache.is_swa_enabled = False + cache.cache_controller = None + cache.metrics_collector = None + cache.tree_core = MagicMock() + cache.token_to_kv_pool_allocator = MagicMock() + + capacity = {"available": 0} + mamba_allocator = MagicMock() + mamba_allocator.schedulable_available_size.side_effect = lambda: capacity[ + "available" + ] + cache.req_to_token_pool = MagicMock(mamba_allocator=mamba_allocator) + + def next_node(component_type, tracker): + return (None, False) if tracker[component_type] >= 3 else (1, True) + + def evict_leaf(_node_id, tracker): + tracker[ComponentType.FULL] += 20 + tracker[ComponentType.MAMBA] += 1 + capacity["available"] += 3 + return None + + cache._evict_device_next_node = MagicMock(side_effect=next_node) + cache._evict_device_leaf = MagicMock(side_effect=evict_leaf) + + result = cache.evict_for_alloc(EvictParams(mamba_num=3)) + + self.assertEqual(capacity["available"], 3) + self.assertEqual(result.num_tokens_evicted, 20) + self.assertEqual(result.mamba_num_evicted, 1) + + def test_common_helper_uses_allocation_aware_entry_point(self): + tree_cache = MagicMock() + tree_cache.is_chunk_cache.return_value = False + tree_cache.token_to_kv_pool_allocator.available_size.return_value = 30 + + evict_from_tree_cache(tree_cache, num_tokens=100) + + tree_cache.evict_for_alloc.assert_called_once_with(EvictParams(num_tokens=70)) + tree_cache.evict.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 90a14b8c6..bb3fe457e 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -343,6 +343,7 @@ def build_fixture( cfg: CacheConfig, *, enable_kv_cache_events: bool = False, + enable_session_radix_cache: bool = False, tree_page_size: Optional[int] = None, mamba_cache_chunk_size: Optional[int] = None, ): @@ -472,6 +473,7 @@ def build_fixture( tree_components=cfg.components, enable_mamba_extra_buffer=cfg.enable_mamba_extra_buffer, enable_kv_cache_events=enable_kv_cache_events, + enable_session_radix_cache=enable_session_radix_cache, eviction_policy=cfg.eviction_policy, is_eagle=cfg.is_eagle, ) @@ -481,6 +483,162 @@ def build_fixture( return cache, allocator, req_to_token_pool +@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA") +class TestUnifiedRadixAllocationEvictionRealComponents(CustomTestCase): + """Allocation targets are observed between real auxiliary-tree steps.""" + + _SHORTFALL = 100 + + def _insert(self, cache, allocator, req_to_token_pool, tokens) -> None: + value = allocator.alloc(len(tokens)) + self.assertIsNotNone(value) + params = InsertParams( + key=RadixKey(array("q", tokens)), + value=value[: len(tokens)], + ) + if cache.supports_mamba(): + req = Req( + rid=f"mamba-{len(tokens)}", + origin_input_text="", + origin_input_ids=array("q"), + sampling_params=SamplingParams(temperature=0, max_new_tokens=1), + ) + req_to_token_pool.alloc([req]) + params.mamba_value = req.mamba_pool_idx.unsqueeze(0) + cache.insert(params) + + def _build_internal_chain(self, component_type, enable_session_radix_cache): + cfg = ( + CacheConfig( + components=(ComponentType.FULL, ComponentType.SWA), + sliding_window_size=128, + ) + if component_type is ComponentType.SWA + else CacheConfig( + components=(ComponentType.FULL, ComponentType.MAMBA), + mamba_cache_size=8, + ) + ) + cache, allocator, req_to_token_pool = build_fixture( + cfg, enable_session_radix_cache=enable_session_radix_cache + ) + for length in (2, 4, 6): + self._insert( + cache, + allocator, + req_to_token_pool, + list(range(1, length + 1)), + ) + + lru = cache.tree_core.lru_lists[component_type] + first = lru.get_lru_no_lock() + second = lru.get_prev_no_lock(first) + leaf = lru.get_prev_no_lock(second) + self.assertNotIn(first, cache.tree_core.evictable_device_leaves) + self.assertNotIn(second, cache.tree_core.evictable_device_leaves) + self.assertIn(leaf, cache.tree_core.evictable_device_leaves) + for node in (first, second, leaf): + self.assertIsNotNone(node.component_data[component_type].value) + self.assertIsNotNone(node.component_data[ComponentType.FULL].value) + return cache, first, second, leaf + + def _evict_for_alloc_after_first_drain(self, cache, component_type): + capacity = {"available": 0} + auxiliary_drains = {"count": 0} + real_available_size = cache._component_available_size + real_free_values = cache._free_values + + def available_size(requested_type): + if requested_type is component_type: + return capacity["available"] + return real_available_size(requested_type) + + def free_values(device_frees, host_frees): + freed_auxiliary = bool(device_frees.get(component_type)) + real_free_values(device_frees, host_frees) + if freed_auxiliary: + auxiliary_drains["count"] += 1 + capacity["available"] = self._SHORTFALL + + params = ( + EvictParams(swa_num_tokens=self._SHORTFALL) + if component_type is ComponentType.SWA + else EvictParams(mamba_num=self._SHORTFALL) + ) + with ( + mock.patch.object( + cache, "_component_available_size", side_effect=available_size + ), + mock.patch.object(cache, "_free_values", side_effect=free_values), + ): + result = cache.evict_for_alloc(params) + return result, auxiliary_drains["count"] + + def test_allocation_target_stops_after_one_internal_tombstone(self): + for component_type in (ComponentType.SWA, ComponentType.MAMBA): + for enable_session_radix_cache in (False, True): + with self.subTest( + component_type=component_type, + enable_session_radix_cache=enable_session_radix_cache, + ): + cache, first, second, leaf = self._build_internal_chain( + component_type, enable_session_radix_cache + ) + first_size = len(first.component_data[component_type].value) + + result, drain_count = self._evict_for_alloc_after_first_drain( + cache, component_type + ) + + self.assertIsNone(first.component_data[component_type].value) + self.assertIsNotNone(second.component_data[component_type].value) + self.assertIsNotNone(leaf.component_data[component_type].value) + self.assertIsNotNone(leaf.component_data[ComponentType.FULL].value) + self.assertEqual(result.num_tokens_evicted, 0) + self.assertEqual(drain_count, 1) + evicted = ( + result.swa_num_tokens_evicted + if component_type is ComponentType.SWA + else result.mamba_num_evicted + ) + self.assertEqual(evicted, first_size) + cache.sanity_check() + + def test_explicit_evict_continues_across_internal_steps(self): + for component_type in (ComponentType.SWA, ComponentType.MAMBA): + for enable_session_radix_cache in (False, True): + with self.subTest( + component_type=component_type, + enable_session_radix_cache=enable_session_radix_cache, + ): + cache, first, second, leaf = self._build_internal_chain( + component_type, enable_session_radix_cache + ) + request_count = sum( + len(node.component_data[component_type].value) + for node in (first, second) + ) + params = ( + EvictParams(swa_num_tokens=request_count) + if component_type is ComponentType.SWA + else EvictParams(mamba_num=request_count) + ) + + result = cache.evict(params) + + self.assertIsNone(first.component_data[component_type].value) + self.assertIsNone(second.component_data[component_type].value) + self.assertIsNotNone(leaf.component_data[component_type].value) + self.assertIsNotNone(leaf.component_data[ComponentType.FULL].value) + evicted = ( + result.swa_num_tokens_evicted + if component_type is ComponentType.SWA + else result.mamba_num_evicted + ) + self.assertEqual(evicted, request_count) + cache.sanity_check() + + class TestUnifiedRadixCacheEagleHiCacheStorageKey(CustomTestCase): cfg = CacheConfig( page_size=4, @@ -5363,10 +5521,12 @@ class UnifiedRadixCacheSuite: "alloc", side_effect=[None, retry_slot], ), - mock.patch.object(cache, "evict", autospec=True) as evict, + mock.patch.object( + cache, "evict_for_alloc", autospec=True + ) as evict_for_alloc, ): prep = comp.prepare_load_back(leaf.id, req=req) - evict.assert_called_once_with(EvictParams(num_tokens=0, mamba_num=1)) + evict_for_alloc.assert_called_once_with(EvictParams(num_tokens=0, mamba_num=1)) self.assertIs(prep.allocated_mamba_slot, retry_slot) self.assertEqual(int(req.mamba_pool_idx), int(retry_slot[0])) @@ -5790,13 +5950,15 @@ class UnifiedRadixCacheSuite: int(swa_xfer.host_indices.numel()), ) - with mock.patch.object(cache, "evict", wraps=cache.evict) as evict_mock: + with mock.patch.object( + cache, "evict_for_alloc", wraps=cache.evict_for_alloc + ) as evict_for_alloc_mock: self.assertTrue(cache.load_back(leaf.id)) # Full pre-eviction must not be triggered by SWA pool pressure. full_pre_evict_calls = [ call - for call in evict_mock.call_args_list + for call in evict_for_alloc_mock.call_args_list if call.args and call.args[0].num_tokens > 0 ] self.assertEqual(full_pre_evict_calls, []) @@ -5807,7 +5969,7 @@ class UnifiedRadixCacheSuite: call.args and call.args[0].num_tokens == 0 and call.args[0].swa_num_tokens > 0 - for call in evict_mock.call_args_list + for call in evict_for_alloc_mock.call_args_list ) ) @@ -7026,9 +7188,13 @@ class TestReturnedValuesDrain(_InsertWalkSuite): cases = [ ( "evict_device_next_node", - lambda: make(EvictDeviceNextNodeResult, node_id=node.id), + lambda: make( + EvictDeviceNextNodeResult, + node_id=node.id, + made_progress=True, + ), lambda: cache._evict_device_next_node(ComponentType.FULL, tracker), - node.id, + (node.id, True), ), ( "evict_device_leaf",