[Unified Cache][4/N]: Add Mooncake backend for external linker (#37205)

Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
huangtingwei
2026-09-01 11:40:06 +08:00
committed by GitHub
co-authored by hzh0425
parent 3b14f37b74
commit b21000aef1
2 changed files with 437 additions and 8 deletions
@@ -0,0 +1,414 @@
from __future__ import annotations
import logging
import threading
from concurrent.futures import Future
from queue import Empty, Queue
import torch
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorageConfig,
PoolName,
PoolTransfer,
)
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
HybridCacheController,
)
from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
resolve_hybrid_device_pool_group,
)
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import UnifiedCacheLinker
from sglang.srt.runtime_context import get_memory, get_model
from sglang.srt.utils import freeze_gc, get_device_module
logger = logging.getLogger(__name__)
device_module = get_device_module()
def _storage_suffix(
*, rank_replicated: bool, tp_rank: int, attn_cp_rank: int, pp_rank: int
) -> str:
parts = []
if not rank_replicated:
parts.append(f"tp{tp_rank}")
parts.extend((f"cp{attn_cp_rank}", f"pp{pp_rank}"))
return "_".join(parts)
class LayerWiseLoadCounter:
"""CPU completion counter compatible with KV pools' layer wait hook."""
def __init__(self, num_layers: int):
self.num_layers = num_layers
self.producer_index = -1
self.consumer_index = -1
self.futures: dict[int, list[Future]] = {}
def update_producer(self) -> int:
self.producer_index += 1
self.futures[self.producer_index] = [Future() for _ in range(self.num_layers)]
return self.producer_index
def set_consumer(self, index: int) -> None:
self.consumer_index = index
def complete(self, index: int, layer: int) -> None:
self.futures[index][layer].set_result(None)
def fail(self, index: int, error: BaseException) -> None:
for future in self.futures.get(index, ()):
if not future.done():
future.set_exception(error)
def wait_until(self, threshold: int) -> None:
index = self.consumer_index
futures = self.futures.get(index)
if futures is None:
return
try:
futures[threshold].result()
except BaseException as error:
raise RuntimeError("Mooncake layer-wise KV load failed.") from error
finally:
if threshold == self.num_layers - 1:
self.futures.pop(index, None)
def reset(self) -> None:
self.producer_index = -1
self.consumer_index = -1
self.futures.clear()
class MooncakeDirectLinker(UnifiedCacheLinker):
def __init__(
self,
server_args,
params: CacheInitParams,
*,
components,
storage=None,
):
self.page_size = params.page_size
kvcache = params.token_to_kv_pool_allocator.get_kvcache()
self.pool_group = resolve_hybrid_device_pool_group(
kvcache=kvcache,
page_size=self.page_size,
params=params,
components=components,
)
self.pools = self.pool_group.entry_map
self.num_layers = self.pool_group.num_layers
tp_rank = 0
tp_size = server_args.tp_size
tp_group = params.attn_tp_cache_group or params.tp_cache_group
if torch.distributed.is_available() and torch.distributed.is_initialized():
tp_rank = torch.distributed.get_rank(group=tp_group)
tp_size = torch.distributed.get_world_size(group=tp_group)
rank_replicated = self.pool_group.rank_replicated
self.offload_owner = not rank_replicated or tp_rank == 0
extra_config, *_ = HybridCacheController.parse_storage_backend_extra_config(
get_memory().hicache_storage_backend_extra_config
)
storage_config = HiCacheStorageConfig(
tp_rank=tp_rank,
tp_size=tp_size,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
attn_cp_rank=params.attn_cp_rank,
attn_cp_size=params.attn_cp_size,
is_mla_model=rank_replicated,
enable_storage_metrics=False,
is_page_first_layout=False,
model_name=get_model().model_path,
extra_config=extra_config,
)
if storage is None:
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
MooncakeStore,
)
self.storage = MooncakeStore(storage_config, mem_pool=None)
else:
self.storage = storage
self.storage.mem_pool_host = self.pool_group
self.storage.registered_pools = self.pools
storage_suffix = _storage_suffix(
rank_replicated=rank_replicated,
tp_rank=tp_rank,
attn_cp_rank=params.attn_cp_rank,
pp_rank=params.pp_rank,
)
self.storage.mla_suffix = storage_suffix
self.storage.mha_suffix = storage_suffix
logger.info(
"Mooncake direct linker storage topology: "
"rank_replicated=%s, tp_rank=%d/%d, offload_owner=%s, suffix=%s",
rank_replicated,
tp_rank,
tp_size,
self.offload_owner,
storage_suffix,
)
self.register_buffers()
self.layer_done_counter = LayerWiseLoadCounter(self.num_layers)
if PoolName.MAMBA in self.pools:
params.req_to_token_pool.register_layer_transfer_counter(
self.layer_done_counter
)
self.pending_loads: dict[str, list[PoolTransfer]] = {}
self.gc_frozen = False
self.load_queue: Queue[
tuple[int, dict[str, list[PoolTransfer]], object] | None
] = Queue()
self.completed_loads: Queue[list[str]] = Queue()
self.offload_queue: Queue[tuple[list[PoolTransfer], int, object] | None] = (
Queue()
)
self.offload_results: Queue[bool] = Queue()
self.stats = {"lookup": 0, "load": 0, "offload": 0}
self.load_thread = threading.Thread(
target=self.load_thread_func,
daemon=True,
name=f"mooncake-load-tp{tp_rank}",
)
self.load_thread.start()
self.offload_thread = threading.Thread(
target=self.offload_thread_func,
daemon=True,
name=f"mooncake-offload-tp{tp_rank}",
)
self.offload_thread.start()
def register_buffers(self) -> None:
seen = set()
for pool in self.pools.values():
for buffer in pool.get_hybrid_pool_buffer():
storage = buffer.untyped_storage()
allocation = (int(storage.data_ptr()), int(storage.nbytes()))
if allocation in seen:
continue
seen.add(allocation)
result = self.storage.store.register_buffer(*allocation)
if result not in (0, None):
raise RuntimeError(
"Failed to register GPU KV buffer with Mooncake, "
f"error code: {result}."
)
def lookup(self, rid: str, transfers: list[PoolTransfer]) -> list[int]:
expanded = self.pool_group.resolve_transfers(transfers)
if not expanded:
return []
kv = next(transfer for transfer in transfers if transfer.name == PoolName.KV)
page_keys = list(kv.keys)
if not page_keys:
return []
result = self.storage.batch_exists_v2(page_keys, expanded)
restorable = result.restorable_prefix_pages or []
self.stats["lookup"] += 1
if restorable:
logger.info(
"Mooncake direct linker lookup hit: rid=%s pages=%d candidates=%d",
rid,
restorable[-1],
len(restorable),
)
return restorable
def load(self, rid: str, transfers: list[PoolTransfer]) -> bool:
# Query establishes a boundary at which every component is restorable;
# insert then removes pages already resident in L1. Loading is therefore
# intentionally partial and may contain only a side pool such as SWA.
expanded = self.pool_group.resolve_transfers(
transfers, allow_partial=True, allow_missing_kv=True
)
if not expanded:
return False
if rid in self.pending_loads:
raise RuntimeError(f"Mooncake load for rid={rid} is already queued.")
self.pending_loads[rid] = expanded
return True
def cancel_queued_load(self, rid: str) -> bool:
return self.pending_loads.pop(rid, None) is not None
def num_completed_loads(self) -> int:
return self.completed_loads.qsize()
def pop_completed_load(self) -> list[str]:
return self.completed_loads.get_nowait()
def freeze_gc_once(self) -> None:
if self.gc_frozen:
return
# Transfer metadata creates many short-lived lists. Keep the mature
# model graph out of cyclic GC scans before load or offload traffic.
freeze_gc("Mooncake direct linker")
self.gc_frozen = True
def start_layer_wise_loading(self) -> int:
if not self.pending_loads:
return -1
self.freeze_gc_once()
pending = self.pending_loads
self.pending_loads = {}
counter_index = self.layer_done_counter.update_producer()
ready_event = device_module.Event()
ready_event.record()
self.load_queue.put((counter_index, pending, ready_event))
self.stats["load"] += len(pending)
return counter_index
def load_thread_func(self) -> None:
while True:
task = self.load_queue.get()
try:
if task is None:
return
counter_index, pending, ready_event = task
try:
ready_event.synchronize()
self.load_layer_wise(counter_index, list(pending.values()))
except BaseException as error:
self.layer_done_counter.fail(counter_index, error)
logger.exception("Mooncake layer-wise load batch failed")
finally:
self.completed_loads.put(list(pending))
finally:
self.load_queue.task_done()
def load_layer_wise(
self, counter_index: int, request_transfers: list[list[PoolTransfer]]
) -> None:
started = []
try:
batches: dict[PoolName, tuple[list[str], list[int]]] = {}
for transfers in request_transfers:
for transfer in transfers:
keys, locations = batches.setdefault(transfer.name, ([], []))
component_keys, _ = self.storage._get_hybrid_page_component_keys(
list(transfer.keys), transfer
)
keys.extend(self.storage._tag_keys(component_keys))
locations.extend(
self.pools[transfer.name].prepare_locations(
transfer.host_indices
)
)
for keys, _ in batches.values():
result = self.storage.store.batch_get_session_start(keys)
if list(result) != [0] * len(keys):
raise RuntimeError(
f"Mooncake get session start failed: keys={len(keys)}, "
f"results={result}"
)
started.append(keys)
for layer in range(self.num_layers):
for name, (keys, locations) in batches.items():
meta = self.pools[name].get_prepared_layer_range_meta(
locations, layer
)
if meta is None:
continue
ptrs, sizes, offsets = meta
result = self.storage.store.batch_get_into_multi_buffer_ranges(
keys,
ptrs,
sizes,
offsets,
)
expected = [sum(item) for item in sizes]
if (
result is None
or isinstance(result, int)
or list(result) != expected
):
raise RuntimeError(
f"Mooncake range get failed for pool={name}, "
f"layer={layer}: transferred={result}, "
f"expected={expected}"
)
self.layer_done_counter.complete(counter_index, layer)
except BaseException as error:
self.layer_done_counter.fail(counter_index, error)
logger.exception("Mooncake layer-wise load batch failed")
finally:
for keys in started:
try:
self.storage.store.batch_get_session_end(keys)
except BaseException as error:
self.layer_done_counter.fail(counter_index, error)
logger.exception("Mooncake layer-wise load session cleanup failed")
def offload(self, transfers: list[PoolTransfer]) -> bool:
expanded = self.pool_group.resolve_transfers(transfers, allow_partial=True)
if not expanded:
return False
self.freeze_gc_once()
if not self.offload_owner:
self.offload_results.put(True)
return True
kv = next(transfer for transfer in transfers if transfer.name == PoolName.KV)
tokens = len(kv.keys) * self.page_size
ready_event = device_module.Event()
ready_event.record()
self.offload_queue.put((expanded, tokens, ready_event))
return True
def offload_thread_func(self) -> None:
while True:
task = self.offload_queue.get()
try:
if task is None:
return
expanded, tokens, ready_event = task
ready_event.synchronize()
results = self.storage.batch_set_v2(expanded)
success = all(all(pool_results) for pool_results in results.values())
if success:
self.stats["offload"] += 1
if self.stats["offload"] == 1:
logger.info("Mooncake direct linker offload: tokens=%d", tokens)
self.offload_results.put(success)
except BaseException:
logger.exception("Mooncake offload failed")
self.offload_results.put(False)
finally:
self.offload_queue.task_done()
def num_completed_offloads(self) -> int:
return self.offload_results.qsize()
def pop_completed_offload(self) -> bool:
return self.offload_results.get_nowait()
def reset(self) -> None:
self.pending_loads.clear()
self.load_queue.join()
self.offload_queue.join()
while True:
try:
self.offload_results.get_nowait()
except Empty:
break
while True:
try:
self.completed_loads.get_nowait()
except Empty:
break
self.layer_done_counter.reset()
def close(self) -> None:
self.reset()
self.load_queue.put(None)
self.offload_queue.put(None)
self.load_thread.join()
self.offload_thread.join()
logger.info("Mooncake direct linker stats: %s", self.stats)
self.storage.close()
@@ -753,7 +753,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
# Mooncake zips object keys with registered buffer pointers.
pool_name = transfer.name
suffixes = []
if pool_name == PoolName.MAMBA:
if pool_name == PoolName.KV:
suffixes = [f"_{self.mla_suffix}_k"]
elif pool_name == PoolName.MAMBA:
# Mamba stores one temporal object plus one object per conv state.
# conv-only models have no ssm state; drop the 0-element temporal
# object (mooncake rejects 0-size puts). get_page_buffer_meta drops
@@ -840,10 +842,14 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
kv_pages = self.batch_exists(keys, extra_info)
hit_count: dict = {PoolName.KV: kv_pages} if kv_pages else {}
final_pages = kv_pages
# Start from every KV prefix and let each pool remove the stop points it
# cannot serve. Collect the whole set, not just its maximum: a
# TRAILING_PAGES pool leaves holes (see PoolTransferResult), and the
# caller has to intersect these sets across ranks.
restorable = list(range(1, kv_pages + 1))
for transfer in pool_transfers or []:
if final_pages == 0:
if not restorable:
break
component_keys, key_multiplier = self._get_hybrid_page_component_keys(
keys, transfer
@@ -861,25 +867,34 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
else:
page_exists = [False] * kv_pages
boundary = 0
pool_restorable = []
if transfer.hit_policy == PoolHitPolicy.ALL_PAGES:
try:
boundary = page_exists.index(False)
except ValueError:
boundary = kv_pages
pool_restorable = list(range(1, boundary + 1))
elif transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES:
# A stop point works when the window ending there is complete,
# so scan every one instead of stopping at the longest.
trailing = max(1, len(transfer.keys) if transfer.keys else 1)
for prefix_len in range(kv_pages, 0, -1):
if all(
page_exists[i]
for i in range(max(0, prefix_len - trailing), prefix_len)
):
boundary = prefix_len
break
pool_restorable.append(prefix_len)
if boundary == 0:
boundary = prefix_len
else:
raise ValueError(f"Unsupported pool hit policy: {transfer.hit_policy}")
if boundary:
hit_count[transfer.name] = boundary
final_pages = min(final_pages, boundary)
pool_restorable_set = set(pool_restorable)
restorable = [p for p in restorable if p in pool_restorable_set]
return PoolTransferResult(final_pages, hit_count)
final_pages = restorable[-1] if restorable else 0
return PoolTransferResult(final_pages, hit_count, restorable)
def _batch_io_v2(self, transfers: List[PoolTransfer], is_set: bool):
# Unified v2 I/O path: each PoolTransfer can expand to one or more
@@ -899,7 +914,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
)
key_strs = self._tag_keys(key_strs)
ptr_list, element_size_list = host_pool.get_page_buffer_meta(host_indices)
if transfer.name == PoolName.DEEPSEEK_V4_C4:
if len(ptr_list) != len(key_strs):
ptr_list, element_size_list = self._pack_multi_buffer_meta(
key_strs, ptr_list, element_size_list
)