Add host tensor allocator for memory_pool_host and support Mooncake standalone storage (#14873)

Signed-off-by: Xingrui Yi <yixingrui@linux.alibaba.com>
Co-authored-by: Teng Ma <sima.mt@alibaba-inc.com>
This commit is contained in:
EkiRui
2025-12-28 16:28:50 -08:00
committed by GitHub
co-authored by Teng Ma
parent 6f9d0a89a0
commit ef4b3c0e96
5 changed files with 164 additions and 27 deletions
+2
View File
@@ -251,6 +251,7 @@ class Envs:
# Mooncake Store # Mooncake Store
SGLANG_HICACHE_MOONCAKE_CONFIG_PATH = EnvStr(None) SGLANG_HICACHE_MOONCAKE_CONFIG_PATH = EnvStr(None)
MOONCAKE_MASTER = EnvStr(None) MOONCAKE_MASTER = EnvStr(None)
MOONCAKE_CLIENT = EnvStr(None)
MOONCAKE_LOCAL_HOSTNAME = EnvStr("localhost") MOONCAKE_LOCAL_HOSTNAME = EnvStr("localhost")
MOONCAKE_TE_META_DATA_SERVER = EnvStr("P2PHANDSHAKE") MOONCAKE_TE_META_DATA_SERVER = EnvStr("P2PHANDSHAKE")
MOONCAKE_GLOBAL_SEGMENT_SIZE = EnvStr("4gb") MOONCAKE_GLOBAL_SEGMENT_SIZE = EnvStr("4gb")
@@ -258,6 +259,7 @@ class Envs:
MOONCAKE_DEVICE = EnvStr("") MOONCAKE_DEVICE = EnvStr("")
MOONCAKE_MASTER_METRICS_PORT = EnvInt(9003) MOONCAKE_MASTER_METRICS_PORT = EnvInt(9003)
MOONCAKE_CHECK_SERVER = EnvBool(False) MOONCAKE_CHECK_SERVER = EnvBool(False)
MOONCAKE_STANDALONE_STORAGE = EnvBool(False)
# AMD & ROCm # AMD & ROCm
SGLANG_USE_AITER = EnvBool(False) SGLANG_USE_AITER = EnvBool(False)
@@ -52,6 +52,7 @@ class HiRadixCache(RadixCache):
server_args.hicache_size, server_args.hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
) )
elif isinstance(self.kv_cache, MLATokenToKVPool): elif isinstance(self.kv_cache, MLATokenToKVPool):
self.token_to_kv_pool_host = MLATokenToKVPoolHost( self.token_to_kv_pool_host = MLATokenToKVPoolHost(
@@ -60,6 +61,7 @@ class HiRadixCache(RadixCache):
server_args.hicache_size, server_args.hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
) )
else: else:
raise ValueError(f"HiRadixCache only supports MHA and MLA yet") raise ValueError(f"HiRadixCache only supports MHA and MLA yet")
@@ -52,17 +52,51 @@ def synchronized(func):
return wrapper return wrapper
class HostTensorAllocator(abc.ABC):
def __init__(self):
"""Initialize the HostTensorAllocator."""
self.dtype = None
self.dims = None
def allocate(self, dims: tuple, dtype: torch.dtype, device: str) -> torch.Tensor:
"""Allocate a tensor of given dims and dtype on the memory."""
self.dtype = dtype
self.dims = dims
tensor = torch.empty(dims, dtype=dtype, device=device)
return tensor
def get_allocator_from_storage(allocator_type):
if allocator_type == "mooncake":
try:
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
MooncakeHostTensorAllocator,
)
return MooncakeHostTensorAllocator()
except ImportError:
logger.warning(
"Mooncake's tensor allocator requires mooncake >= 0.3.8. "
"Please upgrade Mooncake by 'pip install mooncake --upgrade'. "
"Fallback to use default allocator."
)
return HostTensorAllocator()
else:
return HostTensorAllocator()
def alloc_with_host_register( def alloc_with_host_register(
dims, dims,
dtype: torch.dtype, dtype: torch.dtype,
device: str, device: str,
pin_memory: bool, pin_memory: bool,
allocator: HostTensorAllocator,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Allocate tensor and register host memory with cudaHostRegister. Allocate tensor and register host memory with cudaHostRegister.
CudaHostRegister only applies when pin_memory=True. CudaHostRegister only applies when pin_memory=True.
""" """
buffer = torch.empty(dims, dtype=dtype, device=device) buffer = allocator.allocate(dims, dtype=dtype, device=device)
if pin_memory: if pin_memory:
torch.cuda.cudart().cudaHostRegister( torch.cuda.cudart().cudaHostRegister(
buffer.data_ptr(), buffer.numel() * buffer.element_size(), 0 buffer.data_ptr(), buffer.numel() * buffer.element_size(), 0
@@ -75,6 +109,7 @@ def alloc_with_pin_memory(
dtype: torch.dtype, dtype: torch.dtype,
device: str, device: str,
pin_memory: bool, pin_memory: bool,
allocator: None,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Allocate tensor using PyTorch's built-in pin_memory flag. Allocate tensor using PyTorch's built-in pin_memory flag.
@@ -102,12 +137,14 @@ class HostKVCache(abc.ABC):
layout: str, layout: str,
pin_memory: bool, pin_memory: bool,
device: str, device: str,
allocator_type: str = "default",
): ):
self.device_pool = device_pool self.device_pool = device_pool
self.page_size = page_size self.page_size = page_size
self.layout = layout self.layout = layout
self.pin_memory = pin_memory self.pin_memory = pin_memory
self.device = device self.device = device
self.allocator = get_allocator_from_storage(allocator_type)
self.dtype = device_pool.store_dtype self.dtype = device_pool.store_dtype
self.size_per_token = self.get_size_per_token() self.size_per_token = self.get_size_per_token()
@@ -239,6 +276,7 @@ class MHATokenToKVPoolHost(HostKVCache):
layout: str, layout: str,
pin_memory: bool = True, pin_memory: bool = True,
device: str = "cpu", device: str = "cpu",
allocator_type: str = "default",
): ):
super().__init__( super().__init__(
device_pool, device_pool,
@@ -248,6 +286,7 @@ class MHATokenToKVPoolHost(HostKVCache):
layout, layout,
pin_memory, pin_memory,
device, device,
allocator_type,
) )
self.element_dim = self.device_pool.head_num * self.device_pool.head_dim self.element_dim = self.device_pool.head_num * self.device_pool.head_dim
self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel( self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel(
@@ -307,7 +346,11 @@ class MHATokenToKVPoolHost(HostKVCache):
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device] alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
buffer = alloc_func( buffer = alloc_func(
dims, dtype=self.dtype, device=self.device, pin_memory=self.pin_memory dims,
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
) )
return buffer return buffer
@@ -645,6 +688,7 @@ class MLATokenToKVPoolHost(HostKVCache):
layout: str, layout: str,
pin_memory: bool = True, pin_memory: bool = True,
device: str = "cpu", device: str = "cpu",
allocator_type: str = "default",
): ):
super().__init__( super().__init__(
device_pool, device_pool,
@@ -654,6 +698,7 @@ class MLATokenToKVPoolHost(HostKVCache):
layout, layout,
pin_memory, pin_memory,
device, device,
allocator_type,
) )
self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)] self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)]
self.data_ptrs = torch.tensor( self.data_ptrs = torch.tensor(
@@ -715,12 +760,14 @@ class MLATokenToKVPoolHost(HostKVCache):
dtype=self.dtype, dtype=self.dtype,
device=self.device, device=self.device,
pin_memory=self.pin_memory, pin_memory=self.pin_memory,
allocator=self.allocator,
) )
self.v_buffer = alloc_func( self.v_buffer = alloc_func(
(*base_dims, self.qk_rope_head_dim), (*base_dims, self.qk_rope_head_dim),
dtype=self.dtype, dtype=self.dtype,
device=self.device, device=self.device,
pin_memory=self.pin_memory, pin_memory=self.pin_memory,
allocator=self.allocator,
) )
# Return k_buffer to preserve original kv_buffer and data_refs init logic, # Return k_buffer to preserve original kv_buffer and data_refs init logic,
# though Ascend doesn't use these parameters. # though Ascend doesn't use these parameters.
@@ -734,7 +781,11 @@ class MLATokenToKVPoolHost(HostKVCache):
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device] alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
buffer = alloc_func( buffer = alloc_func(
dims, dtype=self.dtype, device=self.device, pin_memory=self.pin_memory dims,
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
) )
return buffer return buffer
@@ -163,7 +163,7 @@ class StorageBackendFactory:
elif backend_name == "nixl": elif backend_name == "nixl":
return backend_class(storage_config) return backend_class(storage_config)
elif backend_name == "mooncake": elif backend_name == "mooncake":
backend = backend_class(storage_config) backend = backend_class(storage_config, mem_pool_host)
return backend return backend
elif backend_name == "aibrix": elif backend_name == "aibrix":
backend = backend_class(storage_config, mem_pool_host) backend = backend_class(storage_config, mem_pool_host)
@@ -1,3 +1,4 @@
import ctypes
import json import json
import logging import logging
import os import os
@@ -15,7 +16,7 @@ from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorageConfig, HiCacheStorageConfig,
HiCacheStorageExtraInfo, HiCacheStorageExtraInfo,
) )
from sglang.srt.mem_cache.memory_pool_host import HostKVCache from sglang.srt.mem_cache.memory_pool_host import HostKVCache, HostTensorAllocator
DEFAULT_LOCAL_BUFFER_SIZE = 16 * 1024 * 1024 # 16 MB DEFAULT_LOCAL_BUFFER_SIZE = 16 * 1024 * 1024 # 16 MB
SETUP_TIMEOUT = 600 # 10min SETUP_TIMEOUT = 600 # 10min
@@ -23,6 +24,41 @@ SETUP_TIMEOUT = 600 # 10min
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class MooncakeHostTensorAllocator(HostTensorAllocator):
def __init__(self):
super().__init__()
from mooncake.store import MooncakeHostMemAllocator
self.allocator = MooncakeHostMemAllocator()
self.ptr = None
def allocate(
self, dims: tuple, dtype: torch.dtype, device: str = "cpu"
) -> torch.Tensor:
"""
Allocates memory using MooncakeHostMemAllocator and wraps it in a PyTorch tensor.
"""
self.dims = dims
self.dtype = dtype
size = 1
for d in dims:
size *= d
size *= torch.tensor([], dtype=self.dtype).element_size()
ptr_int = self.allocator.alloc(size)
self.ptr = ptr_int
c_type = ctypes.c_byte * size
c_array = c_type.from_address(ptr_int)
tensor = torch.frombuffer(c_array, dtype=torch.uint8, count=size)
if dtype != torch.uint8:
element_size = torch.tensor([], dtype=dtype).element_size()
assert size % element_size == 0, "Size must be divisible by element size"
tensor = tensor.view(dtype)
return tensor.view(dims)
def _parse_global_segment_size(value) -> int: def _parse_global_segment_size(value) -> int:
if isinstance(value, int): if isinstance(value, int):
return value return value
@@ -49,6 +85,8 @@ class MooncakeStoreConfig:
master_server_address: str master_server_address: str
master_metrics_port: int master_metrics_port: int
check_server: bool check_server: bool
standalone_storage: bool
client_server_address: str
@staticmethod @staticmethod
def from_file() -> "MooncakeStoreConfig": def from_file() -> "MooncakeStoreConfig":
@@ -64,8 +102,13 @@ class MooncakeStoreConfig:
except Exception as e: except Exception as e:
raise RuntimeError(f"Failed to load config from {file_path}: {str(e)}") raise RuntimeError(f"Failed to load config from {file_path}: {str(e)}")
if "master_server_address" not in config: if (
raise ValueError("master_server_address is required in config file") "master_server_address" not in config
and "client_server_address" not in config
):
raise ValueError(
"Either master_server_address or client_server_address is required in config file"
)
return MooncakeStoreConfig( return MooncakeStoreConfig(
local_hostname=config.get( local_hostname=config.get(
@@ -81,11 +124,19 @@ class MooncakeStoreConfig:
), ),
protocol=config.get("protocol", envs.MOONCAKE_PROTOCOL.default), protocol=config.get("protocol", envs.MOONCAKE_PROTOCOL.default),
device_name=config.get("device_name", envs.MOONCAKE_DEVICE.default), device_name=config.get("device_name", envs.MOONCAKE_DEVICE.default),
master_server_address=config.get("master_server_address"), master_server_address=config.get(
"master_server_address", envs.MOONCAKE_MASTER.default
),
master_metrics_port=config.get( master_metrics_port=config.get(
"master_metrics_port", envs.MOONCAKE_MASTER_METRICS_PORT.default "master_metrics_port", envs.MOONCAKE_MASTER_METRICS_PORT.default
), ),
check_server=config.get("check_server", envs.MOONCAKE_CHECK_SERVER.default), check_server=config.get("check_server", envs.MOONCAKE_CHECK_SERVER.default),
standalone_storage=config.get(
"standalone_storage", envs.MOONCAKE_STANDALONE_STORAGE.default
),
client_server_address=config.get(
"client_server_address", envs.MOONCAKE_CLIENT.default
),
) )
@staticmethod @staticmethod
@@ -97,8 +148,10 @@ class MooncakeStoreConfig:
export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE" export MOONCAKE_TE_META_DATA_SERVER="P2PHANDSHAKE"
""" """
# other required environment variables... # other required environment variables...
if not envs.MOONCAKE_MASTER.is_set(): if not envs.MOONCAKE_MASTER.is_set() and not envs.MOONCAKE_CLIENT.is_set():
raise ValueError("The environment variable 'MOONCAKE_MASTER' is not set.") raise ValueError(
"Either the environment variable 'MOONCAKE_MASTER' or 'MOONCAKE_CLIENT' is not set."
)
# Special handling for local_hostname: try MOONCAKE_LOCAL_HOSTNAME first, # Special handling for local_hostname: try MOONCAKE_LOCAL_HOSTNAME first,
# then fall back to LOCAL_HOSTNAME if not set. # then fall back to LOCAL_HOSTNAME if not set.
@@ -121,13 +174,20 @@ class MooncakeStoreConfig:
master_server_address=envs.MOONCAKE_MASTER.get(), master_server_address=envs.MOONCAKE_MASTER.get(),
master_metrics_port=envs.MOONCAKE_MASTER_METRICS_PORT.get(), master_metrics_port=envs.MOONCAKE_MASTER_METRICS_PORT.get(),
check_server=envs.MOONCAKE_CHECK_SERVER.get(), check_server=envs.MOONCAKE_CHECK_SERVER.get(),
standalone_storage=envs.MOONCAKE_STANDALONE_STORAGE.get(),
client_server_address=envs.MOONCAKE_CLIENT.get(),
) )
@staticmethod @staticmethod
def load_from_extra_config(extra_config: dict) -> "MooncakeStoreConfig": def load_from_extra_config(extra_config: dict) -> "MooncakeStoreConfig":
"""Load config from extra_config dictionary.""" """Load config from extra_config dictionary."""
if "master_server_address" not in extra_config: if (
raise ValueError("master_server_address is required in extra_config") "master_server_address" not in extra_config
and "client_server_address" not in extra_config
):
raise ValueError(
"Either master_server_address or client_server_address is required in extra_config"
)
return MooncakeStoreConfig( return MooncakeStoreConfig(
local_hostname=extra_config.get( local_hostname=extra_config.get(
@@ -143,19 +203,29 @@ class MooncakeStoreConfig:
), ),
protocol=extra_config.get("protocol", envs.MOONCAKE_PROTOCOL.default), protocol=extra_config.get("protocol", envs.MOONCAKE_PROTOCOL.default),
device_name=extra_config.get("device_name", envs.MOONCAKE_DEVICE.default), device_name=extra_config.get("device_name", envs.MOONCAKE_DEVICE.default),
master_server_address=extra_config["master_server_address"], master_server_address=extra_config.get(
"master_server_address", envs.MOONCAKE_MASTER.default
),
master_metrics_port=extra_config.get( master_metrics_port=extra_config.get(
"master_metrics_port", envs.MOONCAKE_MASTER_METRICS_PORT.default "master_metrics_port", envs.MOONCAKE_MASTER_METRICS_PORT.default
), ),
check_server=extra_config.get( check_server=extra_config.get(
"check_server", envs.MOONCAKE_CHECK_SERVER.default "check_server", envs.MOONCAKE_CHECK_SERVER.default
), ),
standalone_storage=extra_config.get(
"standalone_storage", envs.MOONCAKE_STANDALONE_STORAGE.default
),
client_server_address=extra_config.get(
"client_server_address", envs.MOONCAKE_CLIENT.default
),
) )
class MooncakeStore(HiCacheStorage): class MooncakeStore(HiCacheStorage):
def __init__(self, storage_config: HiCacheStorageConfig = None): def __init__(
self, storage_config: HiCacheStorageConfig = None, mem_pool: HostKVCache = None
):
try: try:
from mooncake.store import MooncakeDistributedStore from mooncake.store import MooncakeDistributedStore
except ImportError as e: except ImportError as e:
@@ -174,9 +244,9 @@ class MooncakeStore(HiCacheStorage):
else None else None
) )
# Load configuration with master_server_address prioritized from extra_config if available # Load configuration with master_server_address prioritized from extra_config if available
if ( if extra_config is not None and (
extra_config is not None extra_config.get("master_server_address") is not None
and extra_config.get("master_server_address") is not None or extra_config.get("client_server_address") is not None
): ):
# Load from extra_config # Load from extra_config
self.config = MooncakeStoreConfig.load_from_extra_config(extra_config) self.config = MooncakeStoreConfig.load_from_extra_config(extra_config)
@@ -226,16 +296,28 @@ class MooncakeStore(HiCacheStorage):
f"Failed to parse device_name as JSON: {device_name}" f"Failed to parse device_name as JSON: {device_name}"
) )
device_name = "" device_name = ""
if self.config.standalone_storage:
ret_code = self.store.setup( if not isinstance(mem_pool.allocator, MooncakeHostTensorAllocator):
self.config.local_hostname, raise RuntimeError(
self.config.metadata_server, "MooncakeStore with standalone_storage=True requires MooncakeHostTensorAllocator. "
per_tp_global_segment_size, "Please set standalone_storage=False "
DEFAULT_LOCAL_BUFFER_SIZE, # Zero copy interface does not need local buffer "or upgrade Mooncake by 'pip install mooncake --upgrade'."
self.config.protocol, )
device_name, ret_code = self.store.setup_dummy(
self.config.master_server_address, mem_pool.size * mem_pool.size_per_token,
) DEFAULT_LOCAL_BUFFER_SIZE, # Zero copy interface does not need local buffer
self.config.client_server_address,
)
else:
ret_code = self.store.setup(
self.config.local_hostname,
self.config.metadata_server,
per_tp_global_segment_size,
DEFAULT_LOCAL_BUFFER_SIZE, # Zero copy interface does not need local buffer
self.config.protocol,
device_name,
self.config.master_server_address,
)
if ret_code: if ret_code:
raise RuntimeError( raise RuntimeError(
f"Failed to setup Mooncake store, error code: {ret_code}" f"Failed to setup Mooncake store, error code: {ret_code}"