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 a88d4ad74..9eb36ee2c 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/README.md +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/README.md @@ -294,7 +294,12 @@ You can enable it in any of the three supported configuration methods: For a comprehensive overview of HiCache-related parameters, please refer to [this document](https://docs.sglang.io/advanced_features/hicache_design.html#related-parameters). -Note that, for `--hicache-mem-layout {layer_first,page_first,page_first_direct}`, which specifies the memory layout for the host memory pool, `page_first` or `page_first_direct` are required if use Mooncake backend. +Note that, for `--hicache-mem-layout {layer_first,page_first,page_first_direct}`, +the regular Mooncake backend path still uses `page_first` or `page_first_direct`. +When HiSparse provides an MLA host KV pool or DeepSeek V4 C4 side pool with +layer-first page metadata, Mooncake Store uses Mooncake's multi-buffer zero-copy +APIs (`batch_put_from_multi_buffers` / `batch_get_into_multi_buffers`) to store +each logical page across its per-layer buffers. ### Distributed Deployment 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 71281bed9..3c1ab73dc 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 @@ -4,6 +4,7 @@ import logging import os import time import uuid +from collections.abc import Sequence from dataclasses import dataclass from typing import Any, List, Optional, Tuple @@ -540,6 +541,17 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): logger.error("An error occurred while loading the configuration: %s", exc) raise + @staticmethod + def _iter_host_pool_buffers(host_pool: HostKVCache): + get_buffers = getattr( + host_pool, + "get_hybrid_pool_buffer", + lambda: [getattr(host_pool, "kv_buffer", None)], + ) + for buf in get_buffers(): + if buf is not None: + yield buf + def check_server(self): master_server_ip = self.config.master_server_address.split(":")[0] segments_url = f"http://{master_server_ip}:{self.config.master_metrics_port}/get_all_segments" @@ -605,15 +617,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): # Hybrid logical anchors only own allocation indices. Their physical # tensors are registered through register_mem_host_pool_v2(). return - assert self.mem_pool_host.layout in [ - "page_first", - "page_first_direct", - "page_head", - "page_first_kv_split", - ], "mooncake store storage backend only support page first, page first direct, page head and page_first_kv_split layout" - buffer = self.mem_pool_host.kv_buffer try: - super().register_buffer(buffer) + for buffer in self._iter_host_pool_buffers(self.mem_pool_host): + super().register_buffer(buffer) except TypeError as err: logger.error("Failed to register buffer to Mooncake Store: %s", err) raise TypeError("Mooncake Store Register Buffer Error.") from err @@ -637,14 +643,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): # Non-anchor pools are either sidecar-specific pools with their own # accessor, or ordinary KV-like host pools used as SWA side pools. - get_buffers = getattr( - host_pool, - "get_hybrid_pool_buffer", - lambda: [getattr(host_pool, "kv_buffer", None)], - ) - for buf in get_buffers(): - if buf is None: - continue + for buf in self._iter_host_pool_buffers(host_pool): super().register_buffer(buf) def _tag_keys(self, keys: List[str]) -> List[str]: @@ -783,11 +782,15 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): assert len(keys) > 0 assert len(keys) == len(host_indices) // page_size - ptr_list, element_size_list = host_pool.get_page_buffer_meta(host_indices) key_strs, key_multiplier = self._get_hybrid_page_component_keys( keys, transfer ) 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: + ptr_list, element_size_list = self._pack_multi_buffer_meta( + key_strs, ptr_list, element_size_list + ) if is_set: exist_result = self._batch_exist(key_strs) @@ -838,13 +841,40 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): assert len(key_list) == len(ptr_list) return key_list, ptr_list, element_size_list + @staticmethod + def _uses_multi_buffer(buffer_ptrs: List[Any]) -> bool: + return bool(buffer_ptrs) and isinstance(buffer_ptrs[0], Sequence) + + @staticmethod + def _pack_multi_buffer_meta( + key_strs: List[str], + ptr_list: List[int], + element_size_list: List[int], + ) -> Tuple[List[Any], List[Any]]: + if len(ptr_list) == len(key_strs): + return ptr_list, element_size_list + + assert len(key_strs) > 0 + assert len(ptr_list) == len(element_size_list) + assert len(ptr_list) % len(key_strs) == 0 + + nbuf = len(ptr_list) // len(key_strs) + return [ptr_list[i : i + nbuf] for i in range(0, len(ptr_list), nbuf)], [ + element_size_list[i : i + nbuf] + for i in range(0, len(element_size_list), nbuf) + ] + def _get_mha_buffer_meta(self, keys, indices): ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) key_list = [] for key_ in keys: key_list.append(f"{key_}_{self.mha_suffix}_k") key_list.append(f"{key_}_{self.mha_suffix}_v") - assert len(key_list) == len(ptr_list) + if len(key_list) != len(ptr_list): + raise RuntimeError( + "Mooncake layer_first multi-buffer is only supported for MLA " + "host KV pool. Use page_first/page_first_direct for MHA." + ) return key_list, ptr_list, element_size_list def _get_mla_buffer_meta(self, keys, indices): @@ -852,6 +882,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): key_list = [] for key_ in keys: key_list.append(f"{key_}_{self.mla_suffix}_k") + ptr_list, element_size_list = self._pack_multi_buffer_meta( + key_list, ptr_list, element_size_list + ) assert len(key_list) == len(ptr_list) return key_list, ptr_list, element_size_list @@ -1136,13 +1169,21 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): self.store.remove_all() def _put_batch_zero_copy_impl( - self, key_strs: List[str], buffer_ptrs: List[int], buffer_sizes: List[int] + 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_put_from_multi_buffers( + key_strs, buffer_ptrs, buffer_sizes + ) return self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes) def _get_batch_zero_copy_impl( - self, key_strs: List[str], buffer_ptrs: List[int], buffer_sizes: List[int] + 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) def _batch_exist(self, key_strs: List[str]) -> List[int]: