Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 90 additions & 0 deletions benchmark/kvcache/benchmark_mooncake_failed_get_cache.py
Original file line number Diff line number Diff line change
@@ -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()
12 changes: 12 additions & 0 deletions python/sglang/srt/mem_cache/storage/mooncake_store/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down
161 changes: 151 additions & 10 deletions python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Comment thread
catyans marked this conversation as resolved.


class MooncakeHostTensorAllocator(HostTensorAllocator):
def __init__(self):
super().__init__()
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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 "
Expand Down Expand Up @@ -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,
Expand All @@ -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
Comment thread
catyans marked this conversation as resolved.

def get_stats(self):
storage_metrics = StorageMetrics()
Expand Down
Loading
Loading