[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). 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 ### Distributed Deployment
@@ -4,6 +4,7 @@ import logging
import os import os
import time import time
import uuid import uuid
from collections.abc import Sequence
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, List, Optional, Tuple 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) logger.error("An error occurred while loading the configuration: %s", exc)
raise 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): def check_server(self):
master_server_ip = self.config.master_server_address.split(":")[0] 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" 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 # Hybrid logical anchors only own allocation indices. Their physical
# tensors are registered through register_mem_host_pool_v2(). # tensors are registered through register_mem_host_pool_v2().
return 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: 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: except TypeError as err:
logger.error("Failed to register buffer to Mooncake Store: %s", err) logger.error("Failed to register buffer to Mooncake Store: %s", err)
raise TypeError("Mooncake Store Register Buffer Error.") from 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 # Non-anchor pools are either sidecar-specific pools with their own
# accessor, or ordinary KV-like host pools used as SWA side pools. # accessor, or ordinary KV-like host pools used as SWA side pools.
get_buffers = getattr( for buf in self._iter_host_pool_buffers(host_pool):
host_pool,
"get_hybrid_pool_buffer",
lambda: [getattr(host_pool, "kv_buffer", None)],
)
for buf in get_buffers():
if buf is None:
continue
super().register_buffer(buf) super().register_buffer(buf)
def _tag_keys(self, keys: List[str]) -> List[str]: def _tag_keys(self, keys: List[str]) -> List[str]:
@@ -783,11 +782,15 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
assert len(keys) > 0 assert len(keys) > 0
assert len(keys) == len(host_indices) // page_size 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( key_strs, key_multiplier = self._get_hybrid_page_component_keys(
keys, transfer keys, transfer
) )
key_strs = self._tag_keys(key_strs) 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: if is_set:
exist_result = self._batch_exist(key_strs) exist_result = self._batch_exist(key_strs)
@@ -838,13 +841,40 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
assert len(key_list) == len(ptr_list) assert len(key_list) == len(ptr_list)
return key_list, ptr_list, element_size_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): def _get_mha_buffer_meta(self, keys, indices):
ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices)
key_list = [] key_list = []
for key_ in keys: for key_ in keys:
key_list.append(f"{key_}_{self.mha_suffix}_k") key_list.append(f"{key_}_{self.mha_suffix}_k")
key_list.append(f"{key_}_{self.mha_suffix}_v") 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 return key_list, ptr_list, element_size_list
def _get_mla_buffer_meta(self, keys, indices): def _get_mla_buffer_meta(self, keys, indices):
@@ -852,6 +882,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
key_list = [] key_list = []
for key_ in keys: for key_ in keys:
key_list.append(f"{key_}_{self.mla_suffix}_k") 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) assert len(key_list) == len(ptr_list)
return key_list, ptr_list, element_size_list return key_list, ptr_list, element_size_list
@@ -1136,13 +1169,21 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
self.store.remove_all() self.store.remove_all()
def _put_batch_zero_copy_impl( 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]: ) -> 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) return self.store.batch_put_from(key_strs, buffer_ptrs, buffer_sizes)
def _get_batch_zero_copy_impl( 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]: ) -> 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) return self.store.batch_get_into(key_strs, buffer_ptrs, buffer_sizes)
def _batch_exist(self, key_strs: List[str]) -> List[int]: def _batch_exist(self, key_strs: List[str]) -> List[int]: