[HiCache][AMD] Add UMBP tiered DRAM + SSD L3 storage backend with hugepage host allocator (#25377)

Co-authored-by: TianDi101 ditian12@amd.com
Co-authored-by: Niko Ma nima@amd.com
Co-authored-by: Wu, Yutong yutong.wu@amd.com
Co-authored-by: figo fizhang@amd.com
Co-authored-by: AMD-yanfeiwang <yanfei.wang@amd.com>
Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
inkcherry
2026-07-01 22:21:37 +08:00
committed by GitHub
co-authored by TianDi101 ditian12@amd.com Niko Ma nima@amd.com Wu, Yutong yutong.wu@amd.com figo fizhang@amd.com AMD-yanfeiwang Zhangheng
parent a0d9791810
commit 13dc5f2dc7
9 changed files with 2106 additions and 1 deletions
@@ -483,7 +483,7 @@ class HiCacheController:
if (
self.storage_backend_type
in ["hf3fs", "mooncake", "eic", "nixl", "simm"]
in ["hf3fs", "mooncake", "eic", "nixl", "simm", "mori"]
) or (
self.storage_backend_type == "dynamic"
and bool(self.storage_config.extra_config.get("interface_v1", 0))
@@ -40,6 +40,20 @@ def get_allocator_from_storage(allocator_type):
"Fallback to use default allocator."
)
return HostTensorAllocator()
elif allocator_type == "mori":
try:
from sglang.srt.mem_cache.storage.umbp.umbp_host_allocator import (
UMBPHostTensorAllocator,
)
return UMBPHostTensorAllocator()
except (ImportError, RuntimeError) as exc:
logger.warning(
"UMBPHostTensorAllocator unavailable (%s). "
"Falling back to torch.empty-based allocator.",
exc,
)
return HostTensorAllocator()
else:
return HostTensorAllocator()
@@ -185,6 +185,8 @@ class StorageBackendFactory:
return backend_class(storage_config, mem_pool_host)
elif backend_name == "simm":
return backend_class(storage_config, mem_pool_host)
elif backend_name == "mori":
return backend_class(storage_config, mem_pool_host)
else:
raise ValueError(f"Unknown built-in backend: {backend_name}")
@@ -229,3 +231,9 @@ StorageBackendFactory.register_backend(
"sglang.srt.mem_cache.storage.simm.hicache_simm",
"HiCacheSiMM",
)
StorageBackendFactory.register_backend(
"mori",
"sglang.srt.mem_cache.storage.umbp.umbp_store",
"UMBPStore",
)
@@ -0,0 +1,142 @@
import ctypes
import logging
import math
import os
from typing import Any, Dict
import torch
from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator
logger = logging.getLogger(__name__)
def _bool_env(name: str, default: bool) -> bool:
raw = os.getenv(name)
if raw is None:
return default
return raw.strip().lower() in ("1", "true", "yes", "on")
def _int_env(name: str, default: int) -> int:
raw = os.getenv(name)
return int(raw) if raw is not None and raw != "" else default
class UMBPHostTensorAllocator(HostTensorAllocator):
"""Allocate the HiCache L2 host tensor from mori's UMBPHostMemAllocator."""
def __init__(self) -> None:
super().__init__()
try:
import mori.umbp as umbp_mod
except ImportError as exc:
raise RuntimeError(
"mori.umbp is not available. Build mori with BUILD_UMBP=ON "
"or fall back to the default torch host allocator."
) from exc
self._mod = umbp_mod
self._allocator = umbp_mod.UMBPHostMemAllocator()
self._use_hugepage = _bool_env("SGLANG_HICACHE_HOST_HUGEPAGE", True)
self._hugepage_size = _int_env(
"SGLANG_HICACHE_HOST_HUGEPAGE_SIZE", 2 * 1024 * 1024
)
self._numa_node = _int_env("SGLANG_HICACHE_HOST_NUMA_NODE", -1)
self._prefault = _bool_env("SGLANG_HICACHE_HOST_PREFAULT", True)
self._handles: Dict[int, Any] = {}
def allocate(
self, dims: tuple, dtype: torch.dtype, device: str = "cpu"
) -> torch.Tensor:
if device != "cpu":
raise ValueError(
"UMBPHostTensorAllocator only supports CPU host memory, "
f"got device={device}"
)
self.dims = dims
self.dtype = dtype
element_size = torch.empty((), dtype=dtype).element_size()
nbytes = math.prod(int(dim) for dim in dims) * element_size
requested_backing = (
self._mod.UMBPHostBufferBacking.AnonymousHugetlb
if self._use_hugepage
else self._mod.UMBPHostBufferBacking.Anonymous
)
handle = self._allocator.alloc(
nbytes,
requested_backing,
self._hugepage_size,
self._numa_node,
self._prefault,
)
if not handle:
raise RuntimeError(
f"UMBPHostMemAllocator.alloc({nbytes} bytes) failed "
f"(requested_backing={requested_backing}, "
f"numa_node={self._numa_node})."
)
self._handles[int(handle.ptr)] = handle
c_array = (ctypes.c_byte * nbytes).from_address(handle.ptr)
tensor = torch.frombuffer(c_array, dtype=torch.uint8, count=nbytes)
if dtype != torch.uint8:
tensor = tensor.view(dtype)
logger.info(
"UMBPHostTensorAllocator: allocated %.2f GB at 0x%x "
"requested_backing=%s actual_backing=%s actual_alignment=%d "
"mapped_size=%d numa_node=%d",
nbytes / 1e9,
handle.ptr,
requested_backing,
handle.actual_backing,
handle.actual_alignment,
handle.mapped_size,
self._numa_node,
)
if (
self._use_hugepage
and handle.actual_backing == self._mod.UMBPHostBufferBacking.Anonymous
):
logger.warning(
"UMBPHostTensorAllocator: requested AnonymousHugetlb backing "
"but kernel demoted to Anonymous (4 KiB pages). Check "
"vm.nr_hugepages and HugePages_Free in /proc/meminfo. "
"Performance and AINIC MR-size benefits will not apply."
)
return tensor.view(dims)
def mapped_size_for(self, ptr: int) -> int:
"""Actual mmap size for the allocation whose base address is *ptr*."""
handles = getattr(self, "_handles", None)
if handles is None:
return 0
h = handles.get(ptr)
return int(h.mapped_size) if h is not None else 0
@property
def mapped_size(self) -> int:
"""Largest mapped_size across all live allocations, or 0."""
handles = getattr(self, "_handles", None)
if not handles:
return 0
return max(int(h.mapped_size) for h in handles.values())
def __del__(self) -> None:
try:
handles = getattr(self, "_handles", None)
allocator = getattr(self, "_allocator", None)
if handles and allocator is not None:
for h in handles.values():
allocator.free(h)
self._handles.clear()
except Exception:
pass
File diff suppressed because it is too large Load Diff
+1
View File
@@ -1923,6 +1923,7 @@ class ServerArgs:
"dynamic",
"eic",
"simm",
"mori",
],
),
] = None
@@ -0,0 +1,202 @@
import builtins
import ctypes
import gc
import importlib
import sys
import types
import unittest
from enum import Enum
from unittest import mock
import torch
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
# These tests stub out mori with a fake in-process module, so they need neither
# a real mori install nor a GPU and run on NVIDIA / CPU CI.
class FakeBacking(Enum):
Anonymous = 0
AnonymousHugetlb = 1
class FakeHandle:
def __init__(
self,
ptr: int,
requested_size: int,
mapped_size: int,
actual_backing: FakeBacking,
actual_alignment: int,
) -> None:
self.ptr = ptr
self.requested_size = requested_size
self.mapped_size = mapped_size
self.actual_backing = actual_backing
self.actual_alignment = actual_alignment
def __bool__(self) -> bool:
return self.ptr is not None
class FakeHostMemAllocator:
def __init__(self) -> None:
self.alloc_calls = []
self.free_calls = []
self._buffers = []
def alloc(
self,
size: int,
backing: FakeBacking,
hugepage_size: int,
numa_node: int,
prefault: bool,
) -> FakeHandle:
buf = (ctypes.c_byte * size)()
self._buffers.append(buf)
handle = FakeHandle(
ptr=ctypes.addressof(buf),
requested_size=size,
mapped_size=size,
actual_backing=backing,
actual_alignment=(
hugepage_size if backing == FakeBacking.AnonymousHugetlb else 4096
),
)
self.alloc_calls.append(
{
"size": size,
"backing": backing,
"hugepage_size": hugepage_size,
"numa_node": numa_node,
"prefault": prefault,
"handle": handle,
}
)
return handle
def free(self, handle: FakeHandle) -> None:
self.free_calls.append(handle)
handle.ptr = None
handle.requested_size = 0
handle.mapped_size = 0
class TestUMBPHostAllocator(unittest.TestCase):
def _save_mori_modules(self):
"""Snapshot and restore sys.modules entries for mori on cleanup."""
saved = {name: sys.modules.get(name) for name in ("mori", "mori.umbp")}
def restore():
for name, value in saved.items():
if value is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = value
self.addCleanup(restore)
def _install_fake_mori(self):
self._save_mori_modules()
fake_umbp = types.ModuleType("mori.umbp")
fake_umbp.UMBPHostBufferBacking = FakeBacking
fake_umbp.UMBPHostBufferHandle = FakeHandle
fake_umbp.UMBPHostMemAllocator = FakeHostMemAllocator
fake_mori = types.ModuleType("mori")
fake_mori.__path__ = []
fake_mori.umbp = fake_umbp
sys.modules["mori"] = fake_mori
sys.modules["mori.umbp"] = fake_umbp
return fake_umbp
def test_umbp_allocator_dispatch_and_tensor_wrap(self):
self._install_fake_mori()
from sglang.srt.mem_cache.memory_pool_host import get_allocator_from_storage
from sglang.srt.mem_cache.storage.umbp.umbp_host_allocator import (
UMBPHostTensorAllocator,
)
allocator = get_allocator_from_storage("mori")
self.assertIsInstance(allocator, UMBPHostTensorAllocator)
tensor = allocator.allocate((2, 3), dtype=torch.float16, device="cpu")
alloc_call = allocator._allocator.alloc_calls[0]
self.assertEqual(tensor.shape, (2, 3))
self.assertEqual(tensor.dtype, torch.float16)
self.assertEqual(tensor.data_ptr(), alloc_call["handle"].ptr)
self.assertEqual(alloc_call["size"], tensor.numel() * tensor.element_size())
self.assertEqual(alloc_call["backing"], FakeBacking.AnonymousHugetlb)
self.assertEqual(alloc_call["hugepage_size"], 2 * 1024 * 1024)
self.assertEqual(alloc_call["numa_node"], -1)
self.assertIs(alloc_call["prefault"], True)
tensor.fill_(3.0)
self.assertEqual(float(tensor[0, 0]), 3.0)
def test_umbp_allocator_del_calls_free_once(self):
self._install_fake_mori()
module = importlib.import_module(
"sglang.srt.mem_cache.storage.umbp.umbp_host_allocator"
)
allocator = module.UMBPHostTensorAllocator()
tensor = allocator.allocate((16,), dtype=torch.uint8, device="cpu")
del tensor
gc.collect()
fake_allocator = allocator._allocator
handles = list(allocator._handles.values())
self.assertEqual(len(handles), 1)
handle = handles[0]
allocator.__del__()
self.assertEqual(len(fake_allocator.free_calls), 1)
self.assertIs(fake_allocator.free_calls[0], handle)
self.assertIsNone(handle.ptr)
self.assertEqual(handle.requested_size, 0)
self.assertEqual(handle.mapped_size, 0)
allocator.__del__()
self.assertEqual(len(fake_allocator.free_calls), 1)
def test_get_allocator_from_storage_umbp_falls_back(self):
self._save_mori_modules()
sys.modules.pop("mori", None)
sys.modules.pop("mori.umbp", None)
real_import = builtins.__import__
def fake_import(name, globals=None, locals=None, fromlist=(), level=0):
if name == "mori" or name.startswith("mori."):
raise ImportError("mori unavailable in test")
return real_import(name, globals, locals, fromlist, level)
from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator
with mock.patch.object(builtins, "__import__", fake_import):
with self.assertLogs(level="WARNING") as cm:
from sglang.srt.mem_cache.memory_pool_host import (
get_allocator_from_storage,
)
allocator = get_allocator_from_storage("mori")
self.assertIs(type(allocator), HostTensorAllocator)
self.assertTrue(
any("UMBPHostTensorAllocator unavailable" in msg for msg in cm.output),
f"missing fallback warning in logs: {cm.output}",
)
if __name__ == "__main__":
unittest.main()
+294
View File
@@ -0,0 +1,294 @@
#!/usr/bin/env python3
"""Unit tests for UMBPStore with mocked HostKVCache."""
import ctypes
import tempfile
import unittest
from dataclasses import dataclass
from typing import Optional
from unittest.mock import MagicMock
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
# UMBPStore wraps mori's UMBP client (AMD/ROCm only). On machines without mori
# (e.g. NVIDIA / CPU CI) the whole TestCase is skipped instead of failing at
# import time, so the CI runner (`python3 <file> -f`) exits cleanly.
try:
import mori.umbp # noqa: F401
HAS_MORI = True
except ImportError:
HAS_MORI = False
@dataclass
class MockStorageConfig:
tp_rank: int = 0
tp_size: int = 1
pp_rank: int = 0
pp_size: int = 1
is_mla_model: bool = False
is_page_first_layout: bool = True
model_name: str = "test-model"
tp_lcm_size: Optional[int] = None
should_split_heads: bool = False
extra_config: Optional[dict] = None
class MockHostKVCache:
"""Mock HostKVCache that simulates page_first layout with real buffers."""
def __init__(self, num_pages=4, page_size=1, element_size=1024):
self.layout = "page_first"
self.page_size = page_size
self.element_size = element_size # bytes per K or V per page
total_bytes = num_pages * 2 * element_size # K+V for each page
self._buffer = (ctypes.c_char * total_bytes)()
self._buffer_ptr = ctypes.addressof(self._buffer)
self.kv_buffer = MagicMock()
self.kv_buffer.data_ptr.return_value = self._buffer_ptr
def get_page_buffer_meta(self, indices):
"""Return (ptr_list, element_size_list) for MHA page_first layout.
For page_first MHA: alternating K, V pointers per page.
"""
ptr_list = []
pages = list(range(0, len(indices), self.page_size))
for page_start in pages:
page_idx = (
indices[page_start] if hasattr(indices, "__getitem__") else page_start
)
# K pointer
k_ptr = self._buffer_ptr + page_idx * 2 * self.element_size
# V pointer
v_ptr = k_ptr + self.element_size
ptr_list.append(k_ptr)
ptr_list.append(v_ptr)
return ptr_list, self.element_size
def fill_page(self, page_idx, k_val, v_val):
"""Fill a page's K and V with specific byte values."""
k_offset = page_idx * 2 * self.element_size
v_offset = k_offset + self.element_size
ctypes.memset(self._buffer_ptr + k_offset, k_val, self.element_size)
ctypes.memset(self._buffer_ptr + v_offset, v_val, self.element_size)
def read_page_k(self, page_idx):
"""Read K data for a page."""
k_offset = page_idx * 2 * self.element_size
return bytes(ctypes.string_at(self._buffer_ptr + k_offset, self.element_size))
def read_page_v(self, page_idx):
"""Read V data for a page."""
v_offset = page_idx * 2 * self.element_size + self.element_size
return bytes(ctypes.string_at(self._buffer_ptr + v_offset, self.element_size))
def make_indices(indices):
"""Create a list that acts like a torch.Tensor of indices."""
return indices
@unittest.skipUnless(HAS_MORI, "mori.umbp not available (AMD/ROCm only)")
class TestUMBPStore(unittest.TestCase):
def test_basic_set_get(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=4, page_size=1, element_size=512)
store.register_mem_pool_host(mem_pool)
# Fill page 0 with data
mem_pool.fill_page(0, ord("A"), ord("B"))
# Set: store page 0 data
keys = ["hash_page_0"]
indices = make_indices([0])
result = store.batch_set_v1(keys, indices)
self.assertEqual(len(result), 1)
self.assertTrue(result[0], f"Set failed: {result}")
# Clear the buffer to prove get actually reads from store
mem_pool.fill_page(0, 0, 0)
# Get: restore page 0 data
result = store.batch_get_v1(keys, indices)
self.assertEqual(len(result), 1)
self.assertTrue(result[0], f"Get failed: {result}")
# Verify data restored
k_data = mem_pool.read_page_k(0)
v_data = mem_pool.read_page_v(0)
self.assertEqual(k_data, bytes([ord("A")] * 512), "K data mismatch")
self.assertEqual(v_data, bytes([ord("B")] * 512), "V data mismatch")
def test_batch_set_get_multiple_pages(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 4 * 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=4, page_size=1, element_size=256)
store.register_mem_pool_host(mem_pool)
# Fill pages with distinct data
for i in range(4):
mem_pool.fill_page(i, ord("A") + i, ord("a") + i)
keys = [f"hash_{i}" for i in range(4)]
indices = make_indices([0, 1, 2, 3])
# Set all 4 pages
set_results = store.batch_set_v1(keys, indices)
self.assertTrue(all(set_results), f"Batch set failed: {set_results}")
# Clear buffer
for i in range(4):
mem_pool.fill_page(i, 0, 0)
# Get all 4 pages
get_results = store.batch_get_v1(keys, indices)
self.assertTrue(all(get_results), f"Batch get failed: {get_results}")
# Verify each page
for i in range(4):
k = mem_pool.read_page_k(i)
v = mem_pool.read_page_v(i)
self.assertEqual(k[0], ord("A") + i, f"Page {i} K mismatch")
self.assertEqual(v[0], ord("a") + i, f"Page {i} V mismatch")
def test_batch_exists(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=4, page_size=1, element_size=256)
store.register_mem_pool_host(mem_pool)
# Store first 2 pages
for i in range(2):
mem_pool.fill_page(i, ord("X"), ord("Y"))
keys_to_set = [f"exists_{i}" for i in range(2)]
indices = make_indices([0, 1])
store.batch_set_v1(keys_to_set, indices)
# Check exists: first 2 exist, 3rd does not
all_keys = [f"exists_{i}" for i in range(3)]
count = store.batch_exists(all_keys)
self.assertEqual(count, 2, f"Expected 2 consecutive, got {count}")
def test_dedup_on_set(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=2, page_size=1, element_size=256)
store.register_mem_pool_host(mem_pool)
mem_pool.fill_page(0, ord("A"), ord("B"))
# Set once
keys = ["dedup_key"]
indices = make_indices([0])
store.batch_set_v1(keys, indices)
# Set again — should succeed (dedup)
mem_pool.fill_page(0, ord("X"), ord("Y")) # Different data
result = store.batch_set_v1(keys, indices)
self.assertTrue(result[0])
# Get should return original data (dedup means second set was skipped)
mem_pool.fill_page(0, 0, 0)
store.batch_get_v1(keys, indices)
k = mem_pool.read_page_k(0)
self.assertEqual(k[0], ord("A"), f"Expected original data 'A', got {chr(k[0])}")
def test_clear(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=2, page_size=1, element_size=256)
store.register_mem_pool_host(mem_pool)
mem_pool.fill_page(0, ord("C"), ord("D"))
store.batch_set_v1(["clear_key"], make_indices([0]))
self.assertTrue(store.exists("clear_key_0_k"))
store.clear()
self.assertFalse(store.exists("clear_key_0_k"))
def test_legacy_interface(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
config = MockStorageConfig(
extra_config={"dram_capacity_bytes": 1024 * 1024, "ssd_enabled": False}
)
store = UMBPStore(config)
# Direct set/get/exists via legacy interface
data = (ctypes.c_char * 256)(*([b"Z"] * 256))
ptr = ctypes.addressof(data)
self.assertTrue(store.set("legacy_key", target_location=ptr, target_sizes=256))
self.assertTrue(store.exists("legacy_key"))
buf = (ctypes.c_char * 256)()
result = store.get(
"legacy_key", target_location=ctypes.addressof(buf), target_sizes=256
)
self.assertIsNotNone(result)
self.assertEqual(buf[0], b"Z")
def test_segmented_layout_basic(self):
from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore
with tempfile.TemporaryDirectory(prefix="umbp_segmented_") as ssd_dir:
config = MockStorageConfig(
extra_config={
"dram_capacity_bytes": 1024 * 1024,
"ssd_enabled": True,
"ssd_storage_dir": ssd_dir,
"ssd_capacity_bytes": 16 * 1024 * 1024,
}
)
store = UMBPStore(config)
mem_pool = MockHostKVCache(num_pages=2, page_size=1, element_size=256)
store.register_mem_pool_host(mem_pool)
mem_pool.fill_page(0, ord("M"), ord("N"))
keys = ["seg_hash_0"]
indices = make_indices([0])
self.assertEqual(store.batch_set_v1(keys, indices), [True])
mem_pool.fill_page(0, 0, 0)
self.assertEqual(store.batch_get_v1(keys, indices), [True])
self.assertEqual(mem_pool.read_page_k(0)[0], ord("M"))
self.assertEqual(mem_pool.read_page_v(0)[0], ord("N"))
store.clear()
if __name__ == "__main__":
unittest.main()