[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:
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
@@ -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
@@ -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
@@ -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()
|
||||
Reference in New Issue
Block a user