feat(hicache): Add shared memory allocator for host KV cache (#29326)

Co-authored-by: Jimmy Shong <69131491+Jiminator@users.noreply.github.com>
Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
Bill Du
2026-07-25 12:05:37 +08:00
committed by GitHub
co-authored by Jimmy Shong Zhangheng
parent b83041c3cc
commit ebcb74abd4
14 changed files with 422 additions and 14 deletions
@@ -18,6 +18,7 @@ from sglang.srt.mem_cache.memory_pool import (
MLATokenToKVPool,
ReqToTokenPool,
)
from sglang.srt.mem_cache.pool_host.common import get_allocator_type
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.server_args import ServerArgs
@@ -54,6 +55,8 @@ class DecodeKVCacheOffloadManager:
self.page_size, (env_stride // self.page_size) * self.page_size
)
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
allocator_type = get_allocator_type(server_args)
if isinstance(kv_cache, MHATokenToKVPool):
self.decode_host_mem_pool = get_mha_host_pool_cls(kv_cache)(
kv_cache,
@@ -61,6 +64,7 @@ class DecodeKVCacheOffloadManager:
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
allocator_type=allocator_type,
)
elif isinstance(kv_cache, MLATokenToKVPool):
self.decode_host_mem_pool = MLATokenToKVPoolHost(
@@ -69,6 +73,7 @@ class DecodeKVCacheOffloadManager:
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
allocator_type=allocator_type,
)
else:
raise ValueError("Unsupported KV cache type for decode offload")
+5 -2
View File
@@ -48,6 +48,7 @@ from sglang.srt.mem_cache.memory_pool import (
MiniMaxSparseKVPool,
MLATokenToKVPool,
)
from sglang.srt.mem_cache.pool_host.common import get_allocator_type
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.mem_cache.radix_cache import (
@@ -80,6 +81,8 @@ class HiRadixCache(RadixCache):
self.page_size = params.page_size
self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache()
allocator_type = get_allocator_type(server_args)
if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
self.kv_cache,
@@ -87,7 +90,7 @@ class HiRadixCache(RadixCache):
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=allocator_type,
)
elif isinstance(self.kv_cache, DSATokenToKVPool):
# Filled by attach_hybrid_dsa_pool_to_hiradix_cache after storage extra_config is parsed.
@@ -102,7 +105,7 @@ class HiRadixCache(RadixCache):
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=allocator_type,
)
else:
raise ValueError("HiRadixCache only supports MHA, MLA, DSA, and MSA models")
@@ -21,6 +21,7 @@ from sglang.srt.mem_cache.memory_pool_host import (
MambaPoolHost,
PoolEntry,
)
from sglang.srt.mem_cache.pool_host.common import get_allocator_type
from sglang.srt.mem_cache.pool_host.mha import (
MHATokenToKOnlyPoolHost,
get_mha_host_pool_cls,
@@ -40,6 +41,10 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _get_allocator_type(server_args: ServerArgs) -> str:
return get_allocator_type(server_args)
def _make_layer_mapper(
layer_mapping: dict[int, int],
transfer_layer_num: int,
@@ -72,7 +77,7 @@ def build_kv_host_pool(
server_args.hicache_size,
page_size,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
**kwargs,
)
@@ -353,7 +358,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=swa_num_host_pages,
slot_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
entries.append(
@@ -379,7 +384,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
c4_indexer_host_pool = DeepSeekV4PagedHostPool(
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER),
@@ -391,7 +396,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
entries.extend(
[
@@ -422,7 +427,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE),
@@ -433,7 +438,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
entries.extend(
[
@@ -465,7 +470,7 @@ def build_deepseek_v4_hicache_stack(
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
)
# C128 state pool is intentionally not registered with hicache.
# page_size=256 % 128 == 0, so state pool is not consumed on load.
@@ -533,7 +538,7 @@ def build_hybrid_mamba_stack(
mamba_pool,
server_args.hicache_ratio,
server_args.hicache_size,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
layout=server_args.hicache_mem_layout,
)
entries = [
@@ -1107,7 +1112,7 @@ class _DsaStrategy(StackStrategy):
full_kv_pool,
kv_host_pool,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
),
prefetch_threshold=prefetch_threshold,
model_name=model_name,
@@ -1545,7 +1550,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
kv,
kv_host_pool,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
allocator_type=_get_allocator_type(server_args),
),
model_name=server_args.served_model_name,
storage_backend_extra_config=extra_config,
@@ -139,6 +139,7 @@ class HostKVCache(abc.ABC):
)
self.kv_buffer = self.init_kv_buffer()
self.fd = getattr(self.allocator, "fd", None)
# A lock for synchronized operations on memory allocation and state transitions.
self.lock = threading.RLock()
@@ -1,11 +1,13 @@
from __future__ import annotations
import json
import logging
import os
from collections import defaultdict
import torch
from sglang.srt.mem_cache.mmap_allocator import alloc_mmap
from sglang.srt.mem_cache.storage.mmap import alloc_mmap
logger = logging.getLogger(__name__)
@@ -25,6 +27,43 @@ class HostTensorAllocator:
return alloc_mmap(dims, dtype)
class ShmHostTensorAllocator(HostTensorAllocator):
def __init__(self):
super().__init__()
self.fds = []
self.mms = []
@property
def fd(self):
return self.fds[0] if self.fds else None
@property
def mm(self):
return self.mms[0] if self.mms else None
def allocate(self, dims: tuple, dtype: torch.dtype, device: str) -> torch.Tensor:
assert (
device == "cpu"
), f"ShmHostTensorAllocator only supports CPU allocations; got device={device!r}"
self.dtype = dtype
self.dims = dims
from sglang.srt.mem_cache.storage.mmap import alloc_shm
tensor, fd, mm = alloc_shm(dims, dtype)
self.fds.append(fd)
self.mms.append(mm)
return tensor
def __del__(self):
for fd in getattr(self, "fds", []):
if fd is not None:
try:
os.close(fd)
except OSError:
pass
self.fds = []
def get_allocator_from_storage(allocator_type):
if allocator_type == "mooncake":
try:
@@ -54,10 +93,30 @@ def get_allocator_from_storage(allocator_type):
exc,
)
return HostTensorAllocator()
elif allocator_type == "shm":
return ShmHostTensorAllocator()
else:
return HostTensorAllocator()
def get_allocator_type(server_args) -> str:
backend = getattr(server_args, "hicache_storage_backend", None)
if backend == "shm":
return "shm"
if backend == "dynamic":
extra_config_str = getattr(
server_args, "hicache_storage_backend_extra_config", None
)
if extra_config_str:
try:
config = json.loads(extra_config_str)
if config.get("allocator") == "shm":
return "shm"
except Exception:
pass
return backend or "default"
def _cuda_host_register(buffer: torch.Tensor) -> None:
cudart = torch.cuda.cudart()
n_bytes = buffer.numel() * buffer.element_size()
@@ -187,6 +187,8 @@ class StorageBackendFactory:
return backend_class(storage_config, mem_pool_host)
elif backend_name == "mori":
return backend_class(storage_config, mem_pool_host)
elif backend_name == "shm":
return backend_class(storage_config, mem_pool_host)
else:
raise ValueError(f"Unknown built-in backend: {backend_name}")
@@ -237,3 +239,9 @@ StorageBackendFactory.register_backend(
"sglang.srt.mem_cache.storage.umbp.umbp_store",
"UMBPStore",
)
StorageBackendFactory.register_backend(
"shm",
"sglang.srt.mem_cache.storage.shm",
"HiCacheShm",
)
@@ -0,0 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to SGLang project
"""Mmap allocator storage backend helpers for SGLang HiCache."""
from .mmap_allocator import alloc_mmap, alloc_shm
__all__ = [
"alloc_mmap",
"alloc_shm",
]
@@ -4,6 +4,7 @@ import logging
import math
import mmap
import os
import uuid
import weakref
import torch
@@ -39,6 +40,7 @@ _MAP_HUGETLB = 0x40000
_MAP_HUGE_2MB = 21 << 26 # 0x1400000
_MAP_HUGE_1GB = 30 << 26 # 0x78000000
_MAP_FAILED = ctypes.c_void_p(-1).value
_MADV_POPULATE_WRITE = getattr(mmap, "MADV_POPULATE_WRITE", 23)
def _alloc_hugepage(n_bytes: int, alloc_bytes: int, extra_flags: int) -> ctypes.Array:
@@ -124,4 +126,76 @@ def alloc_mmap(dims: tuple, dtype: torch.dtype) -> torch.Tensor:
flags=mmap.MAP_SHARED | mmap.MAP_ANONYMOUS | _MAP_POPULATE,
prot=mmap.PROT_READ | mmap.PROT_WRITE,
)
try:
# MADV_POPULATE_WRITE guarantees pages are populated and writable,
# throwing an error on failure (e.g. out of memory).
mm.madvise(_MADV_POPULATE_WRITE)
except OSError:
# Fall back to MAP_POPULATE if MADV_POPULATE_WRITE is not supported (<5.14 kernel).
pass
return torch.frombuffer(mm, dtype=dtype, count=math.prod(dims)).reshape(dims)
def alloc_shm(dims: tuple, dtype: torch.dtype) -> tuple[torch.Tensor, int, mmap.mmap]:
"""Allocate a host tensor via shared memory (/dev/shm).
Returns a tuple of (tensor, fd, mm).
The caller is responsible for keeping the fd open if they need to share it,
and closing it when they are done.
"""
hugepage_size = (envs.SGLANG_HUGEPAGE_SIZE.get() or "").strip().upper()
n_bytes = math.prod(dims) * torch.empty([], dtype=dtype).element_size()
# Note: hugepages are not directly supported with /dev/shm mmap files
# without mounting hugetlbfs there, so we fall back to plain page size.
if hugepage_size != "":
logger.warning(
"Hugepages are not supported with SHM allocator. "
"Falling back to plain page-size mmap."
)
page_size = mmap.PAGESIZE
alloc_bytes = math.ceil(n_bytes / page_size) * page_size
# Create an anonymous shared memory file descriptor via memfd_create
fd = None
try:
# MFD_CLOEXEC is standard on Linux 3.17+
fd = os.memfd_create(
f"sglang_host_pool_{uuid.uuid4().hex}",
flags=getattr(os, "MFD_CLOEXEC", 1),
)
except (AttributeError, OSError):
# Fallback to creating a file in /dev/shm if memfd_create is not supported
shm_path = f"/dev/shm/sglang_host_pool_{uuid.uuid4().hex}.mmap"
try:
fd = os.open(shm_path, os.O_CREAT | os.O_RDWR | os.O_TRUNC, 0o600)
try:
os.unlink(shm_path)
except OSError:
pass
except Exception as e:
raise OSError(f"Failed to create shm file: {e}")
try:
os.ftruncate(fd, alloc_bytes)
mm = mmap.mmap(
fd,
alloc_bytes,
flags=mmap.MAP_SHARED | _MAP_POPULATE,
prot=mmap.PROT_READ | mmap.PROT_WRITE,
)
try:
# MADV_POPULATE_WRITE guarantees pages are populated and writable,
# throwing an error on failure (e.g. out of memory).
mm.madvise(_MADV_POPULATE_WRITE)
except OSError:
# Fall back to MAP_POPULATE if MADV_POPULATE_WRITE is not supported (<5.14 kernel).
pass
except Exception as e:
if fd is not None:
os.close(fd)
raise e
tensor = torch.frombuffer(mm, dtype=dtype, count=math.prod(dims)).reshape(dims)
return tensor, fd, mm
@@ -18,8 +18,8 @@ from sglang.srt.mem_cache.hicache_storage import (
PoolTransfer,
PoolTransferResult,
)
from sglang.srt.mem_cache.mmap_allocator import alloc_mmap
from sglang.srt.mem_cache.pool_host import HostKVCache
from sglang.srt.mem_cache.storage.mmap import alloc_mmap
from sglang.srt.mem_cache.storage.nixl.nixl_cleaner import HiCacheL3Cleaner
from .nixl_registry import NixlRegistry
@@ -0,0 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to SGLang project
"""Shared memory storage backend for SGLang HiCache."""
from .hicache_shm import HiCacheShm
__all__ = ["HiCacheShm"]
@@ -0,0 +1,96 @@
import logging
from typing import Any, List, Optional
import torch
from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorage,
HiCacheStorageConfig,
HiCacheStorageExtraInfo,
PoolTransfer,
PoolTransferResult,
)
logger = logging.getLogger(__name__)
class HiCacheShm(HiCacheStorage):
"""
Dummy storage backend for shared memory allocator.
Since shm is a local allocator, there's no actual storage transfer needed.
"""
def __init__(
self, storage_config: HiCacheStorageConfig, mem_pool_host: Optional[Any] = None
):
pass
def get(
self,
key: str,
target_location: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> torch.Tensor | None:
return None
def batch_get(
self,
keys: List[str],
target_locations: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> List[torch.Tensor | None]:
return [None] * len(keys)
def set(
self,
key: str,
value: Optional[Any] = None,
target_location: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> bool:
return True
def batch_set(
self,
keys: List[str],
values: Optional[Any] = None,
target_locations: Optional[Any] = None,
target_sizes: Optional[Any] = None,
) -> bool:
return True
def exists(self, key: str) -> bool:
return False
def batch_exists_v2(
self,
keys: List[str],
pool_transfers: Optional[List[PoolTransfer]] = None,
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> PoolTransferResult:
return PoolTransferResult(0, {})
def batch_get_v2(
self,
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict[str, List[bool]]:
results = {}
for transfer in transfers:
keys = transfer.keys or []
results[transfer.name] = [False] * len(keys)
return results
def batch_set_v2(
self,
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict[str, List[bool]]:
results = {}
for transfer in transfers:
keys = transfer.keys or []
results[transfer.name] = [True] * len(keys)
return results
def clear(self) -> bool:
return True
+1
View File
@@ -2587,6 +2587,7 @@ class ServerArgs:
"eic",
"simm",
"mori",
"shm",
],
),
NS("memory"),