From b21000aef1c337945f420dbf32385525bb7d6323 Mon Sep 17 00:00:00 2001 From: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:40:06 +0800 Subject: [PATCH] [Unified Cache][4/N]: Add Mooncake backend for external linker (#37205) Co-authored-by: hzh0425 --- .../mooncake_store/mooncake_direct_linker.py | 414 ++++++++++++++++++ .../storage/mooncake_store/mooncake_store.py | 31 +- 2 files changed, 437 insertions(+), 8 deletions(-) create mode 100644 python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py new file mode 100644 index 000000000..b69509981 --- /dev/null +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py @@ -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() 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 efb8edbd8..a9362a1c0 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 @@ -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 )