[unified-memory] Stop eviction when shared allocation capacity is sufficient (#33091)
Co-authored-by: seokwoosong <seokwoosong@users.noreply.github.com>
This commit is contained in:
@@ -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()
|
||||
@@ -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, "
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user