diff --git a/benchmark/kvcache/benchmark_mooncake_failed_get_cache.py b/benchmark/kvcache/benchmark_mooncake_failed_get_cache.py new file mode 100755 index 000000000000..d942ad914dda --- /dev/null +++ b/benchmark/kvcache/benchmark_mooncake_failed_get_cache.py @@ -0,0 +1,90 @@ +#!/usr/bin/env python3 + +import argparse +import json +import time + +from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import ( + MooncakeStore, + _FailedGetCache, +) + + +class DelayedStaleBackend: + def __init__(self, get_delay_seconds: float): + self.get_delay_seconds = get_delay_seconds + self.exists_calls = 0 + self.get_calls = 0 + + def batch_is_exist(self, keys): + self.exists_calls += 1 + return [1] * len(keys) + + def batch_get_into(self, keys, buffer_ptrs, buffer_sizes): + self.get_calls += 1 + time.sleep(self.get_delay_seconds) + return [-5] * len(keys) + + +def run_case(iterations: int, get_delay_seconds: float, ttl_seconds: float): + backend = DelayedStaleBackend(get_delay_seconds) + store = MooncakeStore.__new__(MooncakeStore) + store.store = backend + store.failed_get_cache = ( + _FailedGetCache(ttl_seconds, max_entries=1024) if ttl_seconds > 0 else None + ) + + started = time.perf_counter() + for _ in range(iterations): + if store._batch_exist(["stale-page"])[0] == 1: + store._get_batch_zero_copy_impl(["stale-page"], [0x1000], [4096]) + elapsed = time.perf_counter() - started + + return ( + { + "iterations": iterations, + "get_delay_ms": get_delay_seconds * 1000, + "ttl_seconds": ttl_seconds, + "elapsed_ms": elapsed * 1000, + "remote_exists_calls": backend.exists_calls, + "remote_get_calls": backend.get_calls, + }, + store, + backend, + ) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--iterations", type=int, default=1000) + parser.add_argument("--get-delay-ms", type=float, default=10.0) + parser.add_argument("--ttl-seconds", type=float, default=1.0) + args = parser.parse_args() + + delay_seconds = args.get_delay_ms / 1000 + baseline, _, _ = run_case(args.iterations, delay_seconds, 0) + enabled, store, backend = run_case(args.iterations, delay_seconds, args.ttl_seconds) + + time.sleep(args.ttl_seconds + 0.05) + exists_before_retry = backend.exists_calls + retry_visible = store._batch_exist(["stale-page"])[0] == 1 + + result = { + "baseline": baseline, + "enabled": enabled, + "elapsed_reduction_pct": 100 + * (baseline["elapsed_ms"] - enabled["elapsed_ms"]) + / baseline["elapsed_ms"], + "get_suppression_pct": 100 + * (baseline["remote_get_calls"] - enabled["remote_get_calls"]) + / baseline["remote_get_calls"], + "ttl_retry": { + "remote_exists_calls_added": backend.exists_calls - exists_before_retry, + "remote_hit_visible_after_expiry": retry_visible, + }, + } + print(json.dumps(result, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/README.md b/python/sglang/srt/mem_cache/storage/mooncake_store/README.md index 419ee5c5f466..85c91452e7ff 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/README.md +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/README.md @@ -305,6 +305,18 @@ python -m sglang.launch_server \ --hicache-storage-backend-extra-config '{"master_server_address": "127.0.0.1:50051", "enable_group_semantics": true}' ``` +**Failed-get negative cache:** + +Mooncake metadata can briefly outlive a failed or terminated data segment. To avoid repeatedly issuing a remote get that has just failed, the backend suppresses `batch_exists` hits for the failed physical key for one second. A successful put immediately clears the entry. + +Configure the TTL and memory bound through `--hicache-storage-backend-extra-config`: + +```bash +--hicache-storage-backend-extra-config '{"failed_get_ttl_seconds": 1.0, "failed_get_cache_max_entries": 65536}' +``` + +Set `failed_get_ttl_seconds` to `0` to disable the cache. The cache is process-local, bounded, and affects performance only: a suppressed remote hit falls back to recomputation, so it cannot expose incomplete KV data. + **HiCache Related Parameters for SGLang Server** For a comprehensive overview of HiCache-related parameters, please refer to [this document](https://docs.sglang.io/advanced_features/hicache_design.html#related-parameters). diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py index dd42e2cac511..382b4e51665a 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py @@ -2,8 +2,10 @@ import json import logging import os +import threading import time import uuid +from collections import OrderedDict from collections.abc import Sequence from dataclasses import dataclass from typing import Any, List, Optional, Tuple @@ -31,6 +33,92 @@ logger = logging.getLogger(__name__) +class _FailedGetCache: + """Bounded TTL cache for physical keys that recently failed to load.""" + + def __init__(self, ttl_seconds: float, max_entries: int): + self.ttl_seconds = ttl_seconds + self.max_entries = max_entries + self.cache: OrderedDict[str, float] = OrderedDict() + self.lock = threading.Lock() + + def add(self, key: str) -> None: + self.add_batch([key]) + + def add_batch(self, keys: Sequence[str]) -> None: + now = time.monotonic() + with self.lock: + for key in keys: + self.cache.pop(key, None) + self.cache[key] = now + while self.cache: + _, oldest = next(iter(self.cache.items())) + if ( + len(self.cache) <= self.max_entries + and now - oldest <= self.ttl_seconds + ): + break + self.cache.popitem(last=False) + + def remove(self, key: str) -> None: + self.remove_batch([key]) + + def remove_batch(self, keys: Sequence[str]) -> None: + with self.lock: + for key in keys: + self.cache.pop(key, None) + + def update_batch( + self, successful_keys: Sequence[str], failed_keys: Sequence[str] + ) -> None: + now = time.monotonic() + with self.lock: + for key in successful_keys: + self.cache.pop(key, None) + for key in failed_keys: + self.cache.pop(key, None) + self.cache[key] = now + while self.cache: + _, oldest = next(iter(self.cache.items())) + if ( + len(self.cache) <= self.max_entries + and now - oldest <= self.ttl_seconds + ): + break + self.cache.popitem(last=False) + + def contains(self, key: str) -> bool: + now = time.monotonic() + with self.lock: + failed_at = self.cache.get(key) + if failed_at is None: + return False + if now - failed_at > self.ttl_seconds: + del self.cache[key] + return False + return True + + def filter_failed(self, keys: Sequence[str]) -> tuple[List[int], List[str]]: + now = time.monotonic() + query_indices = [] + query_keys = [] + with self.lock: + for i, key in enumerate(keys): + failed_at = self.cache.get(key) + if failed_at is None: + query_indices.append(i) + query_keys.append(key) + elif now - failed_at > self.ttl_seconds: + del self.cache[key] + query_indices.append(i) + query_keys.append(key) + return query_indices, query_keys + + def clear(self) -> None: + with self.lock: + self.cache.clear() + + class MooncakeHostTensorAllocator(HostTensorAllocator): def __init__(self): super().__init__() @@ -314,7 +402,6 @@ def register_buffer(self, tensor: torch.Tensor): class MooncakeStore(HiCacheStorage, MooncakeBaseStore): - @staticmethod def _standalone_required_bytes(mem_pool: Any) -> int: """Compute total bytes of host buffers that must be visible to the real client. @@ -388,6 +475,23 @@ def __init__( and self._supports_group_ids and self._replicate_config_cls is not None ) + failed_get_ttl = float( + extra_config.get("failed_get_ttl_seconds", 1.0) if extra_config else 1.0 + ) + if failed_get_ttl < 0: + raise ValueError("failed_get_ttl_seconds must be non-negative") + failed_get_cache_max_entries = int( + extra_config.get("failed_get_cache_max_entries", 65536) + if extra_config + else 65536 + ) + if failed_get_cache_max_entries <= 0: + raise ValueError("failed_get_cache_max_entries must be positive") + self.failed_get_cache = ( + _FailedGetCache(failed_get_ttl, failed_get_cache_max_entries) + if failed_get_ttl > 0 + else None + ) if self.enable_group_semantics and not self._supports_group_ids: logger.warning( "Mooncake group semantics is enabled, but the installed " @@ -1248,6 +1352,8 @@ def close(self): def clear(self) -> None: self.store.remove_all() + if self.failed_get_cache is not None: + self.failed_get_cache.clear() def _put_batch_zero_copy_impl( self, @@ -1268,27 +1374,62 @@ def _put_batch_zero_copy_impl( if self._uses_multi_buffer(buffer_ptrs): config = config or self._replicate_config_cls() - return self.store.batch_put_from_multi_buffers( + results = self.store.batch_put_from_multi_buffers( key_strs, buffer_ptrs, buffer_sizes, config ) elif config is not None: - return self.store.batch_put_from( + results = self.store.batch_put_from( key_strs, buffer_ptrs, buffer_sizes, config ) else: - return self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes) + results = self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes) + + if self.failed_get_cache is not None: + successful_keys = [ + key for key, result in zip(key_strs, results) if result >= 0 + ] + self.failed_get_cache.remove_batch(successful_keys) + return results def _get_batch_zero_copy_impl( self, key_strs: List[str], buffer_ptrs: List[Any], buffer_sizes: List[Any] ) -> List[int]: - if self._uses_multi_buffer(buffer_ptrs): - return self.store.batch_get_into_multi_buffers( - key_strs, buffer_ptrs, buffer_sizes - ) - return self.store.batch_get_into(key_strs, buffer_ptrs, buffer_sizes) + try: + if self._uses_multi_buffer(buffer_ptrs): + results = self.store.batch_get_into_multi_buffers( + key_strs, buffer_ptrs, buffer_sizes + ) + else: + results = self.store.batch_get_into(key_strs, buffer_ptrs, buffer_sizes) + except Exception: + if self.failed_get_cache is not None: + self.failed_get_cache.add_batch(key_strs) + raise + + if self.failed_get_cache is not None: + successful_keys = [] + failed_keys = [] + for key, result in zip(key_strs, results): + if result <= 0: + failed_keys.append(key) + else: + successful_keys.append(key) + self.failed_get_cache.update_batch(successful_keys, failed_keys) + return results def _batch_exist(self, key_strs: List[str]) -> List[int]: - return self.store.batch_is_exist(key_strs) + if self.failed_get_cache is None: + return self.store.batch_is_exist(key_strs) + + results = [0] * len(key_strs) + query_indices, query_keys = self.failed_get_cache.filter_failed(key_strs) + if not query_indices: + return results + + query_results = self.store.batch_is_exist(query_keys) + for i, result in zip(query_indices, query_results): + results[i] = result + return results def get_stats(self): storage_metrics = StorageMetrics() diff --git a/test/registered/unit/mem_cache/test_mooncake_failed_get_cache.py b/test/registered/unit/mem_cache/test_mooncake_failed_get_cache.py new file mode 100644 index 000000000000..6d350a571707 --- /dev/null +++ b/test/registered/unit/mem_cache/test_mooncake_failed_get_cache.py @@ -0,0 +1,112 @@ +import sys +from unittest.mock import MagicMock + +import pytest + +from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import ( + MooncakeStore, + _FailedGetCache, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +def _make_store(ttl_seconds: float = 5.0) -> MooncakeStore: + store = MooncakeStore.__new__(MooncakeStore) + store.store = MagicMock() + store.failed_get_cache = _FailedGetCache(ttl_seconds, max_entries=128) + return store + + +def test_failed_get_is_negative_cached_for_exists(): + store = _make_store() + store.store.batch_get_into.return_value = [16, -5] + + assert store._get_batch_zero_copy_impl( + ["good", "stale"], [0x1000, 0x2000], [16, 16] + ) == [16, -5] + + store.store.batch_is_exist.return_value = [1] + assert store._batch_exist(["good", "stale"]) == [1, 0] + store.store.batch_is_exist.assert_called_once_with(["good"]) + + +def test_get_exception_negative_caches_all_attempted_keys(): + store = _make_store() + store.store.batch_get_into.side_effect = RuntimeError("transfer timeout") + + with pytest.raises(RuntimeError, match="transfer timeout"): + store._get_batch_zero_copy_impl( + ["stale-a", "stale-b"], [0x1000, 0x2000], [16, 16] + ) + + assert store._batch_exist(["stale-a", "stale-b"]) == [0, 0] + store.store.batch_is_exist.assert_not_called() + + +def test_successful_put_clears_negative_cache_entry(): + store = _make_store() + store.failed_get_cache.add("restored") + store._use_group_semantics = False + store._replicate_config_cls = None + store.store.batch_put_from.return_value = [0] + + assert store._put_batch_zero_copy_impl(["restored"], [0x1000], [16]) == [0] + + store.store.batch_is_exist.return_value = [1] + assert store._batch_exist(["restored"]) == [1] + store.store.batch_is_exist.assert_called_once_with(["restored"]) + + +def test_failed_put_keeps_negative_cache_entry(): + store = _make_store() + store.failed_get_cache.add("still-stale") + store._use_group_semantics = False + store._replicate_config_cls = None + store.store.batch_put_from.return_value = [-5] + + assert store._put_batch_zero_copy_impl(["still-stale"], [0x1000], [16]) == [-5] + assert store._batch_exist(["still-stale"]) == [0] + store.store.batch_is_exist.assert_not_called() + + +def test_failed_get_cache_entry_expires(monkeypatch): + now = 100.0 + monkeypatch.setattr( + "sglang.srt.mem_cache.storage.mooncake_store.mooncake_store.time.monotonic", + lambda: now, + ) + store = _make_store(ttl_seconds=1.0) + store.failed_get_cache.add("recovered") + + now = 101.01 + store.store.batch_is_exist.return_value = [1] + assert store._batch_exist(["recovered"]) == [1] + store.store.batch_is_exist.assert_called_once_with(["recovered"]) + + +def test_clear_removes_failed_get_entries(): + store = _make_store() + store.failed_get_cache.add("old") + + store.clear() + + store.store.remove_all.assert_called_once_with() + store.store.batch_is_exist.return_value = [1] + assert store._batch_exist(["old"]) == [1] + + +def test_failed_get_cache_is_bounded(): + cache = _FailedGetCache(ttl_seconds=5.0, max_entries=2) + cache.add("oldest") + cache.add("middle") + cache.add("newest") + + assert not cache.contains("oldest") + assert cache.contains("middle") + assert cache.contains("newest") + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"]))