[HiSparse & HiCache]Support mooncake store layer first layout (#27454)
This commit is contained in:
@@ -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]:
|
||||||
|
|||||||
Reference in New Issue
Block a user