[HiSparse & HiCache]Support mooncake store layer first layout (#27454)

This commit is contained in:
huangtingwei
2026-06-07 12:46:42 +08:00
committed by GitHub
parent 52f221cce0
commit 857ecb2dbc
2 changed files with 67 additions and 21 deletions
@@ -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
@@ -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]: