From 13dc5f2dc789e7012a7fd4266e368940a93b7a52 Mon Sep 17 00:00:00 2001 From: inkcherry Date: Wed, 1 Jul 2026 22:21:37 +0800 Subject: [PATCH] [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 Co-authored-by: Zhangheng --- .../sglang/srt/managers/cache_controller.py | 2 +- .../sglang/srt/mem_cache/pool_host/common.py | 14 + .../srt/mem_cache/storage/backend_factory.py | 8 + .../srt/mem_cache/storage/umbp/__init__.py | 0 .../storage/umbp/umbp_host_allocator.py | 142 ++ .../srt/mem_cache/storage/umbp/umbp_store.py | 1444 +++++++++++++++++ python/sglang/srt/server_args.py | 1 + .../mem_cache/test_umbp_host_allocator.py | 202 +++ .../unit/mem_cache/test_umbp_store.py | 294 ++++ 9 files changed, 2106 insertions(+), 1 deletion(-) create mode 100644 python/sglang/srt/mem_cache/storage/umbp/__init__.py create mode 100644 python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py create mode 100644 python/sglang/srt/mem_cache/storage/umbp/umbp_store.py create mode 100644 test/registered/unit/mem_cache/test_umbp_host_allocator.py create mode 100755 test/registered/unit/mem_cache/test_umbp_store.py diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 7505f5fe7..ab1347c0b 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -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)) diff --git a/python/sglang/srt/mem_cache/pool_host/common.py b/python/sglang/srt/mem_cache/pool_host/common.py index db17af43a..07a2cdfb9 100644 --- a/python/sglang/srt/mem_cache/pool_host/common.py +++ b/python/sglang/srt/mem_cache/pool_host/common.py @@ -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() diff --git a/python/sglang/srt/mem_cache/storage/backend_factory.py b/python/sglang/srt/mem_cache/storage/backend_factory.py index 1fe731520..093ac86f1 100644 --- a/python/sglang/srt/mem_cache/storage/backend_factory.py +++ b/python/sglang/srt/mem_cache/storage/backend_factory.py @@ -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", +) diff --git a/python/sglang/srt/mem_cache/storage/umbp/__init__.py b/python/sglang/srt/mem_cache/storage/umbp/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py b/python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py new file mode 100644 index 000000000..fef8ecb3d --- /dev/null +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py @@ -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 diff --git a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py new file mode 100644 index 000000000..5910b75ea --- /dev/null +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py @@ -0,0 +1,1444 @@ +"""UMBPStore — HiCache L3 storage backend using UMBP (local DRAM + SSD). + +Follows the same pattern as MooncakeStore: +- Zero-copy v1 interface (batch_get_v1 / batch_set_v1) +- Uses mem_pool_host.get_page_buffer_meta() for pointer/size extraction +- Key suffix generation per TP rank / PP rank +""" + +import logging +import os +import socket +import threading +from typing import Any, List, Optional + +import torch + +from sglang.srt.mem_cache.hicache_storage import ( + HiCacheStorage, + HiCacheStorageConfig, + HiCacheStorageExtraInfo, +) +from sglang.srt.mem_cache.memory_pool_host import HostKVCache + +logger = logging.getLogger(__name__) + + +def _import_umbp_client(): + """Import UMBPClient from mori.umbp (requires mori built with BUILD_UMBP=ON).""" + import mori.umbp as umbp_mod + + UMBPClient = umbp_mod.UMBPClient + UMBPConfig = umbp_mod.UMBPConfig + UMBPRole = umbp_mod.UMBPRole + UMBPIoBackend = getattr(umbp_mod, "UMBPIoBackend", None) + UMBPDurabilityMode = getattr(umbp_mod, "UMBPDurabilityMode", None) + UMBPDistributedConfig = getattr(umbp_mod, "UMBPDistributedConfig", None) + + return ( + UMBPClient, + UMBPConfig, + UMBPRole, + UMBPIoBackend, + UMBPDurabilityMode, + UMBPDistributedConfig, + ) + + +def _optional_env_int(name: str) -> Optional[int]: + value = os.getenv(name) + return int(value) if value is not None else None + + +def _optional_env_str(name: str) -> Optional[str]: + value = os.getenv(name) + return value if value is not None and value != "" else None + + +_TRUE_STRINGS = frozenset({"1", "true", "yes", "on"}) +_FALSE_STRINGS = frozenset({"0", "false", "no", "off"}) + + +def _strict_bool(value: Any, key: str) -> bool: + """Strict boolean parse; raises rather than silently inverting (bool("false") is True).""" + if isinstance(value, bool): + return value + if isinstance(value, int): + if value in (0, 1): + return bool(value) + elif isinstance(value, str): + norm = value.strip().lower() + if norm in _TRUE_STRINGS: + return True + if norm in _FALSE_STRINGS: + return False + raise ValueError( + f"extra_config[{key!r}] must be a boolean-like value " + f"(true/false, 1/0, yes/no, on/off), got {value!r}" + ) + + +def _cast_like(current: Any, value: Any, key: str) -> Any: + """Cast ``value`` to match the type of an existing config attribute.""" + if isinstance(current, bool): + return _strict_bool(value, key) + if isinstance(current, int): + return int(value) + if isinstance(current, float): + return float(value) + return str(value) + + +def _default_node_address() -> str: + try: + return socket.gethostbyname(socket.gethostname()) + except Exception: + return "127.0.0.1" + + +def _select_rank_config_value( + value: Any, + rank_index: int, + field_name: str, + cast_type, + auto_increment_scalar: bool = False, +): + if value is None: + raise ValueError(f"{field_name} must not be None") + + candidates = value + if isinstance(value, str) and "," in value: + candidates = [item.strip() for item in value.split(",") if item.strip()] + + if isinstance(candidates, (list, tuple)): + if not candidates: + raise ValueError(f"{field_name} must not be empty") + if rank_index >= len(candidates): + raise ValueError( + f"{field_name} has {len(candidates)} entries, but rank_index={rank_index}" + ) + return cast_type(candidates[rank_index]) + + selected = cast_type(candidates) + if auto_increment_scalar: + selected = cast_type(selected + rank_index) + return selected + + +# extra_config is an explicit allow-list grouped by the scope in which each key +# takes effect (distributed mode is enabled by master_address). Advanced SPDK +# knobs outside this list go through the "spdk_passthrough" escape hatch. +_COMMON_EXTRA_KEYS = frozenset( + { + "dram_capacity_bytes", + "ssd_enabled", + "ssd_storage_dir", + "ssd_capacity_bytes", + "ssd_segment_size_bytes", + "ssd_queue_depth", + "ssd_io_backend", + "ssd_durability_mode", + "ssd_backend", + "ssd_high_watermark", + "ssd_low_watermark", + "ssd_copy_queue_depth", + "ssd_copy_worker_threads", + "spdk_nvme_pci_addr", + "spdk_proxy_shm_name", + "spdk_proxy_startup_timeout_ms", + "spdk_proxy_bin", + "spdk_proxy_tenant_id", + "spdk_proxy_tenant_id_base", + "spdk_proxy_tenant_quota_bytes", + "spdk_proxy_max_channels", + "spdk_proxy_data_per_channel_mb", + "spdk_proxy_auto_start", + "spdk_proxy_idle_exit_timeout_ms", + "spdk_proxy_allow_borrow", + "spdk_proxy_reserved_shared_bytes", + "spdk_passthrough", + "kv_events_subscriber", + "kv_events_endpoint", + "kv_events_topic", + } +) + +# No-op in distributed mode (PeerSsdManager/SsdCopyPipeline have no equivalent). +_STANDALONE_ONLY_EXTRA_KEYS = frozenset( + { + "ssd_copy_async", + "ssd_copy_batch_max_ops", + "eviction_policy", + "eviction_candidate_window", + "auto_promote_on_read", + } +) + +_DISTRIBUTED_ONLY_EXTRA_KEYS = frozenset( + { + "master_address", + "node_address", + "node_id", + "auto_heartbeat", + "io_engine_host", + "io_engine_port", + "staging_buffer_size", + "ssd_staging_buffer_size", + "ssd_staging_buffer_slots", + "peer_service_port", + "cache_remote_fetches", + "dram_page_size", + "disable_zero_copy_register", + } +) + +_KNOWN_EXTRA_KEYS = ( + _COMMON_EXTRA_KEYS | _STANDALONE_ONLY_EXTRA_KEYS | _DISTRIBUTED_ONLY_EXTRA_KEYS +) + + +def _warn_extra_config_scope(extra: dict, distributed_enabled: bool) -> None: + """Warn about unknown keys and keys set in a mode where they are no-op.""" + for key in extra: + if key not in _KNOWN_EXTRA_KEYS: + logger.warning( + "UMBPStore: unknown extra_config key %r is ignored. Check for a " + "typo; advanced SPDK knobs go through extra_config['spdk_passthrough'].", + key, + ) + if distributed_enabled: + for key in _STANDALONE_ONLY_EXTRA_KEYS: + if key in extra: + logger.warning( + "UMBPStore: extra_config[%r] is a standalone-only knob and has " + "no effect in distributed mode (master_address set).", + key, + ) + else: + for key in _DISTRIBUTED_ONLY_EXTRA_KEYS: + if key in extra and key != "master_address": + logger.warning( + "UMBPStore: extra_config[%r] only applies in distributed mode " + "(master_address) and is ignored in standalone/local mode.", + key, + ) + + +class UMBPStore(HiCacheStorage): + """Local DRAM+SSD storage backend for HiCache L3 caching. + + Compatible with the zero-copy v1 interface used by CacheController. + """ + + def __init__( + self, + storage_config: HiCacheStorageConfig = None, + mem_pool_host: HostKVCache = None, + ): + ( + UMBPClient, + UMBPConfig, + UMBPRole, + UMBPIoBackend, + UMBPDurabilityMode, + UMBPDistributedConfig, + ) = _import_umbp_client() + + if storage_config is not None: + self.is_mla_backend = storage_config.is_mla_model + self.local_rank = storage_config.tp_rank + self.pp_rank = storage_config.pp_rank + self.pp_size = storage_config.pp_size + self.tp_size = storage_config.tp_size + else: + self.is_mla_backend = False + self.local_rank = 0 + self.pp_rank = 0 + self.pp_size = 1 + self.tp_size = 1 + + cfg = UMBPConfig.from_environment() + # UMBPStore owns role selection explicitly. Do not inherit LOCAL_RANK / + # UMBP_ROLE-based multi-process defaults from mori here, otherwise + # ordinary multi-rank sglang runs can accidentally become follower-only + # and skip writes. + cfg.role = UMBPRole.Standalone + extra = getattr(storage_config, "extra_config", None) or {} + explicit_tenant_id = ( + os.getenv("UMBP_SPDK_PROXY_TENANT_ID") is not None + or "spdk_proxy_tenant_id" in extra + ) + tenant_id_base = ( + int(extra["spdk_proxy_tenant_id_base"]) + if "spdk_proxy_tenant_id_base" in extra + else _optional_env_int("UMBP_SPDK_PROXY_TENANT_ID_BASE") + ) + dp_rank_hint = _optional_env_int("SGLANG_DP_RANK") + dp_size_hint = _optional_env_int("SGLANG_DP_SIZE") + local_rank_hint = _optional_env_int("LOCAL_RANK") + + if dp_rank_hint is None: + try: + from sglang.srt.layers.dp_attention import ( + get_attention_dp_rank, + get_attention_dp_size, + is_dp_attention_enabled, + ) + + if is_dp_attention_enabled(): + dp_rank_hint = get_attention_dp_rank() + dp_size_hint = get_attention_dp_size() + except (ImportError, AssertionError): + pass + + if local_rank_hint is not None: + unique_rank = local_rank_hint + else: + base_rank = dp_rank_hint if dp_rank_hint is not None else 0 + unique_rank = ((base_rank * max(self.pp_size, 1)) + self.pp_rank) * max( + self.tp_size, 1 + ) + self.local_rank + + # Load settings from extra_config if available + if "dram_capacity_bytes" in extra: + cfg.dram.capacity_bytes = int(extra["dram_capacity_bytes"]) + if "ssd_enabled" in extra: + cfg.ssd.enabled = _strict_bool(extra["ssd_enabled"], "ssd_enabled") + if "ssd_storage_dir" in extra: + cfg.ssd.storage_dir = str(extra["ssd_storage_dir"]) + if "ssd_capacity_bytes" in extra: + cfg.ssd.capacity_bytes = int(extra["ssd_capacity_bytes"]) + if "ssd_copy_async" in extra: + cfg.copy_pipeline.async_enabled = _strict_bool( + extra["ssd_copy_async"], "ssd_copy_async" + ) + if "ssd_copy_queue_depth" in extra: + cfg.copy_pipeline.queue_depth = int(extra["ssd_copy_queue_depth"]) + if "ssd_segment_size_bytes" in extra: + cfg.ssd.segment_size_bytes = int(extra["ssd_segment_size_bytes"]) + if "ssd_copy_batch_max_ops" in extra: + cfg.copy_pipeline.batch_max_ops = int(extra["ssd_copy_batch_max_ops"]) + if "ssd_queue_depth" in extra: + cfg.ssd.io.queue_depth = int(extra["ssd_queue_depth"]) + if "ssd_copy_worker_threads" in extra: + cfg.copy_pipeline.worker_threads = int(extra["ssd_copy_worker_threads"]) + if "auto_promote_on_read" in extra: + cfg.eviction.auto_promote_on_read = _strict_bool( + extra["auto_promote_on_read"], "auto_promote_on_read" + ) + if "eviction_policy" in extra: + cfg.eviction.policy = str(extra["eviction_policy"]) + if "eviction_candidate_window" in extra: + cfg.eviction.candidate_window = int(extra["eviction_candidate_window"]) + if "ssd_io_backend" in extra: + backend = str(extra["ssd_io_backend"]).strip().lower() + if backend not in ("posix", "io_uring"): + raise ValueError( + "extra_config['ssd_io_backend'] must be one of: posix, io_uring" + ) + if UMBPIoBackend is not None: + cfg.ssd.io.backend = ( + UMBPIoBackend.Posix if backend == "posix" else UMBPIoBackend.IoUring + ) + if "ssd_durability_mode" in extra: + # Validate even when the enum is unavailable (older mori / mocks). + durability = str(extra["ssd_durability_mode"]).strip().lower() + if durability not in ("strict", "sync", "relaxed", "async"): + raise ValueError( + "extra_config['ssd_durability_mode'] must be one of: " + "strict, sync, relaxed, async" + ) + if UMBPDurabilityMode is not None: + if durability in ("strict", "sync"): + cfg.ssd.durability.mode = UMBPDurabilityMode.Strict + else: + cfg.ssd.durability.mode = UMBPDurabilityMode.Relaxed + if "ssd_backend" in extra: + ssd_backend = str(extra["ssd_backend"]).strip().lower() + if ssd_backend not in ("file", "spdk", "spdk_proxy"): + raise ValueError( + "extra_config['ssd_backend'] must be one of: " + "file, spdk, spdk_proxy" + ) + cfg.ssd.ssd_backend = ssd_backend + if "spdk_nvme_pci_addr" in extra: + cfg.ssd.spdk_nvme_pci_addr = str(extra["spdk_nvme_pci_addr"]) + if "spdk_proxy_shm_name" in extra: + cfg.ssd.spdk_proxy_shm_name = str(extra["spdk_proxy_shm_name"]) + if "spdk_proxy_startup_timeout_ms" in extra: + cfg.ssd.spdk_proxy_startup_timeout_ms = int( + extra["spdk_proxy_startup_timeout_ms"] + ) + if "spdk_proxy_bin" in extra: + cfg.ssd.spdk_proxy_bin = str(extra["spdk_proxy_bin"]) + if "spdk_proxy_tenant_id" in extra: + cfg.ssd.spdk_proxy_tenant_id = int(extra["spdk_proxy_tenant_id"]) + if "spdk_proxy_tenant_quota_bytes" in extra: + cfg.ssd.spdk_proxy_tenant_quota_bytes = int( + extra["spdk_proxy_tenant_quota_bytes"] + ) + if "spdk_proxy_max_channels" in extra: + cfg.ssd.spdk_proxy_max_channels = int(extra["spdk_proxy_max_channels"]) + if "spdk_proxy_data_per_channel_mb" in extra: + cfg.ssd.spdk_proxy_data_per_channel_mb = int( + extra["spdk_proxy_data_per_channel_mb"] + ) + if "spdk_proxy_auto_start" in extra: + cfg.ssd.spdk_proxy_auto_start = _strict_bool( + extra["spdk_proxy_auto_start"], "spdk_proxy_auto_start" + ) + if "spdk_proxy_idle_exit_timeout_ms" in extra: + cfg.ssd.spdk_proxy_idle_exit_timeout_ms = int( + extra["spdk_proxy_idle_exit_timeout_ms"] + ) + if "spdk_proxy_allow_borrow" in extra: + cfg.ssd.spdk_proxy_allow_borrow = _strict_bool( + extra["spdk_proxy_allow_borrow"], "spdk_proxy_allow_borrow" + ) + if "spdk_proxy_reserved_shared_bytes" in extra: + cfg.ssd.spdk_proxy_reserved_shared_bytes = int( + extra["spdk_proxy_reserved_shared_bytes"] + ) + + # Expert escape hatch for advanced SPDK knobs outside the stable + # extra_config surface, forwarded to UMBPSsdConfig as-is. + if "spdk_passthrough" in extra: + overrides = extra["spdk_passthrough"] + if not isinstance(overrides, dict): + raise ValueError( + "extra_config['spdk_passthrough'] must be a dict of " + "spdk_* field overrides" + ) + applied = [] + for field_name, field_value in overrides.items(): + if not field_name.startswith("spdk_") or not hasattr( + cfg.ssd, field_name + ): + raise ValueError( + f"spdk_passthrough: unknown SSD config field {field_name!r} " + "(must be an existing spdk_* field on UMBPSsdConfig)" + ) + current = getattr(cfg.ssd, field_name) + setattr( + cfg.ssd, + field_name, + _cast_like(current, field_value, field_name), + ) + applied.append(field_name) + if applied: + logger.warning( + "UMBPStore: using spdk_passthrough expert escape hatch for " + "fields %s; these are advanced backend knobs forwarded directly " + "to UMBPSsdConfig and are not part of the stable extra_config " + "surface.", + applied, + ) + + if "ssd_high_watermark" in extra and hasattr(cfg.ssd, "high_watermark"): + cfg.ssd.high_watermark = float(extra["ssd_high_watermark"]) + if "ssd_low_watermark" in extra and hasattr(cfg.ssd, "low_watermark"): + cfg.ssd.low_watermark = float(extra["ssd_low_watermark"]) + + # Operator-controlled escape hatch for hosts whose RDMA NIC cannot + # register a single memory region as large as the full host KV buffer + # (e.g. AINIC has a per-MR size cap). When set, skip the one-shot + # register_memory() call in register_mem_pool_host() and stay on the + # staging-buffer fallback path (each transfer copies through a + # staging_buffer_size-bounded MR that the IO engine pre-registers). + disable_zero_copy_register = extra.get( + "disable_zero_copy_register", + _optional_env_str("UMBP_DISABLE_ZERO_COPY_REGISTER"), + ) + self._disable_zero_copy_register = ( + _strict_bool(disable_zero_copy_register, "disable_zero_copy_register") + if disable_zero_copy_register is not None + else False + ) + + master_address = extra.get( + "master_address", _optional_env_str("UMBP_MASTER_ADDRESS") + ) + + _warn_extra_config_scope(extra, distributed_enabled=bool(master_address)) + if master_address and UMBPDistributedConfig is not None: + dist_cfg = UMBPDistributedConfig() + dist_cfg.master_config.master_address = str(master_address) + + if "ssd_copy_worker_threads" not in extra: + cfg.copy_pipeline.worker_threads = 1 + + node_address = extra.get( + "node_address", _optional_env_str("UMBP_NODE_ADDRESS") + ) + if node_address is None: + node_address = _default_node_address() + else: + node_address = _select_rank_config_value( + node_address, + unique_rank, + "node_address", + str, + ) + dist_cfg.master_config.node_address = node_address + + node_id = extra.get("node_id", _optional_env_str("UMBP_NODE_ID")) + if node_id is None: + dist_cfg.master_config.node_id = ( + f"{node_address}:dp{dp_rank_hint if dp_rank_hint is not None else 0}" + f":pp{self.pp_rank}:tp{self.local_rank}" + ) + else: + dist_cfg.master_config.node_id = _select_rank_config_value( + node_id, + unique_rank, + "node_id", + str, + ) + + if "auto_heartbeat" in extra: + dist_cfg.master_config.auto_heartbeat = _strict_bool( + extra["auto_heartbeat"], "auto_heartbeat" + ) + + io_engine_host = extra.get( + "io_engine_host", _optional_env_str("UMBP_IO_ENGINE_HOST") + ) + if io_engine_host is None: + io_engine_host = node_address + else: + io_engine_host = _select_rank_config_value( + io_engine_host, + unique_rank, + "io_engine_host", + str, + ) + dist_cfg.io_engine.host = io_engine_host + + io_engine_port = extra.get( + "io_engine_port", _optional_env_str("UMBP_IO_ENGINE_PORT") + ) + if io_engine_port is not None: + dist_cfg.io_engine.port = _select_rank_config_value( + io_engine_port, + unique_rank, + "io_engine_port", + int, + auto_increment_scalar=True, + ) + + if "staging_buffer_size" in extra: + dist_cfg.staging_buffer_size = int(extra["staging_buffer_size"]) + + if "ssd_staging_buffer_size" in extra and hasattr( + dist_cfg, "ssd_staging_buffer_size" + ): + dist_cfg.ssd_staging_buffer_size = int(extra["ssd_staging_buffer_size"]) + if "ssd_staging_buffer_slots" in extra and hasattr( + dist_cfg, "ssd_staging_buffer_slots" + ): + dist_cfg.ssd_staging_buffer_slots = int( + extra["ssd_staging_buffer_slots"] + ) + + peer_service_port = extra.get( + "peer_service_port", _optional_env_str("UMBP_PEER_SERVICE_PORT") + ) + if peer_service_port is not None: + dist_cfg.peer_service_port = _select_rank_config_value( + peer_service_port, + unique_rank, + "peer_service_port", + int, + auto_increment_scalar=True, + ) + + cache_remote_fetches = extra.get( + "cache_remote_fetches", + _optional_env_str("UMBP_CACHE_REMOTE_FETCHES"), + ) + if cache_remote_fetches is not None: + dist_cfg.cache_remote_fetches = _strict_bool( + cache_remote_fetches, "cache_remote_fetches" + ) + + # Auto-compute master's PageBitmapAllocator page_size so every + # UMBPStore Put/Get maps to exactly one master page (no partial + # tail, 1 RDMA per page). Resolution order: + # 1. extra_config["dram_page_size"] — explicit operator override + # (escape hatch for debugging / forced experiments). + # 2. derived from mem_pool_host (the normal production path). + # 3. left at 0 when neither source is available; mori's + # UMBPDistributedConfig.dram_page_size defaults to 0, which + # delegates to the master-side ClientRegistryConfig + # .default_dram_page_size (2 MiB by default). The + # partial-tail safety net in PoolClient handles any + # size mismatch. + page_byte_size = None + if "dram_page_size" in extra: + page_byte_size = int(extra["dram_page_size"]) + elif mem_pool_host is not None: + # Probe element_size from the same buffer-meta helper that + # batch_preprocess will actually use; this matches per-call + # Put/Get size byte-for-byte for MHA / MHA-split / MLA / NSA + # without per-case formulas (NSA in particular: get_ksize_per_token + # would over-count by the indexer buffer that is never put to UMBP). + dummy = torch.zeros(mem_pool_host.page_size, dtype=torch.int64) + if self.is_mla_backend: + _, esz = mem_pool_host.get_page_buffer_meta(dummy) + elif storage_config is not None and getattr( + storage_config, "should_split_heads", False + ): + sf = storage_config.tp_lcm_size // storage_config.tp_size + _, esz = mem_pool_host.get_split_heads_page_buffer_meta(dummy, sf) + else: + _, esz = mem_pool_host.get_page_buffer_meta(dummy) + page_byte_size = int(esz[0]) if esz else 0 + + if ( + page_byte_size is not None + and page_byte_size > 0 + and hasattr(dist_cfg, "dram_page_size") + ): + dist_cfg.dram_page_size = int(page_byte_size) + logger.info( + "UMBPStore: setting master dram_page_size=%d " + "(ksize_per_token=%s × page_size=%s%s)", + dist_cfg.dram_page_size, + ( + mem_pool_host.get_ksize_per_token() + if mem_pool_host is not None + else "n/a" + ), + (mem_pool_host.page_size if mem_pool_host is not None else "n/a"), + ( + f" / split_factor={storage_config.tp_lcm_size // storage_config.tp_size}" + if ( + mem_pool_host is not None + and storage_config is not None + and getattr(storage_config, "should_split_heads", False) + ) + else "" + ), + ) + + cfg.distributed = dist_cfg + logger.info( + "UMBPStore distributed mode: master=%s, node_id=%s, node_addr=%s, " + "io=%s:%s, peer_port=%s", + dist_cfg.master_config.master_address, + dist_cfg.master_config.node_id, + dist_cfg.master_config.node_address, + dist_cfg.io_engine.host, + dist_cfg.io_engine.port, + dist_cfg.peer_service_port, + ) + + self.storage_config = storage_config + + # MLA + TP > 1: shared SSD mode (standalone only). + # In distributed mode every rank is a peer of the master-led pool; we + # must NOT short-circuit followers (would leave their DRAM pool empty + # while the master still routes keys to them, causing Get misses). + self.is_mla_follower = False + tp_size = self.tp_size + use_spdk = cfg.ssd.ssd_backend in ("spdk", "spdk_proxy") + distributed_enabled = cfg.distributed is not None + if not distributed_enabled and self.is_mla_backend and tp_size > 1: + cfg.ssd.enabled = True + if self.local_rank == 0: + # Leader: copy every DRAM write to shared SSD. + cfg.role = UMBPRole.SharedSSDLeader + else: + # Follower: read-only access. + cfg.role = UMBPRole.SharedSSDFollower + self.is_mla_follower = True + # SPDK: follower must use the proxy path rather than direct + # SpdkSsdTier. Give a longer startup timeout so followers can + # wait for the shared proxy service to become READY. + if use_spdk: + cfg.ssd.ssd_backend = "spdk_proxy" + if cfg.ssd.spdk_proxy_startup_timeout_ms < 60000: + cfg.ssd.spdk_proxy_startup_timeout_ms = 60000 + logger.info( + "UMBPStore MLA+TP>1: rank=%d, role=%s, ssd_backend=%s, shared_ssd=%s", + self.local_rank, + "leader" if self.local_rank == 0 else "follower", + cfg.ssd.ssd_backend, + cfg.ssd.storage_dir, + ) + + try: + from sglang.srt.layers.dp_attention import ( + get_attention_dp_rank, + get_attention_dp_size, + is_dp_attention_enabled, + ) + + if is_dp_attention_enabled(): + dp_rank = get_attention_dp_rank() + dp_size = get_attention_dp_size() + dp_rank_hint = dp_rank + dp_size_hint = dp_size + if cfg.ssd.enabled: + if cfg.ssd.ssd_backend in ("spdk", "spdk_proxy"): + # DP + SPDK must always use the proxy service path. + # Direct SpdkSsdTier is single-process and cannot + # provide tenant isolation across DP ranks. + cfg.ssd.ssd_backend = "spdk_proxy" + if cfg.ssd.spdk_proxy_startup_timeout_ms < 60000: + cfg.ssd.spdk_proxy_startup_timeout_ms = 60000 + if tenant_id_base is not None: + cfg.ssd.spdk_proxy_tenant_id = tenant_id_base + dp_rank + elif not explicit_tenant_id: + cfg.ssd.spdk_proxy_tenant_id = dp_rank + elif dp_size > 1: + logger.warning( + "UMBPStore DP isolation: using explicit fixed tenant_id=%s " + "with dp_size=%d; all DP groups will share one tenant " + "unless you set spdk_proxy_tenant_id_base", + cfg.ssd.spdk_proxy_tenant_id, + dp_size, + ) + if cfg.ssd.spdk_proxy_tenant_quota_bytes <= 0 and dp_size > 1: + # Reserve 5% headroom for offset allocator bin + # rounding (small-float bins round up each + # allocation by up to ~12.5%). + safe_cap = int(cfg.ssd.capacity_bytes * 0.95) + cfg.ssd.spdk_proxy_tenant_quota_bytes = max( + 1, safe_cap // dp_size + ) + # Validate: total tenant quotas must fit within SSD + # capacity after allocator rounding. + if dp_size > 1: + total_quota = ( + cfg.ssd.spdk_proxy_tenant_quota_bytes * dp_size + ) + if total_quota > cfg.ssd.capacity_bytes: + old_quota = cfg.ssd.spdk_proxy_tenant_quota_bytes + safe_cap = int(cfg.ssd.capacity_bytes * 0.95) + cfg.ssd.spdk_proxy_tenant_quota_bytes = max( + 1, safe_cap // dp_size + ) + logger.warning( + "UMBPStore: tenant_quota_bytes=%d × dp_size=%d = %d " + "exceeds ssd_capacity=%d. Reduced to %d to " + "avoid SPDK proxy NO_SPACE. Consider " + "increasing UMBP_SSD_BYTES.", + old_quota, + dp_size, + total_quota, + cfg.ssd.capacity_bytes, + cfg.ssd.spdk_proxy_tenant_quota_bytes, + ) + logger.info( + "UMBPStore DP isolation: dp_rank=%d, dp_size=%d, tenant_id=%s, tenant_quota_bytes=%s", + dp_rank, + dp_size, + getattr(cfg.ssd, "spdk_proxy_tenant_id", "n/a"), + getattr(cfg.ssd, "spdk_proxy_tenant_quota_bytes", "n/a"), + ) + else: + cfg.ssd.storage_dir = f"{cfg.ssd.storage_dir}/dp{dp_rank}" + logger.info( + "UMBPStore DP isolation: dp_rank=%d, dp_size=%d, ssd_dir=%s", + dp_rank, + dp_size, + cfg.ssd.storage_dir, + ) + except (ImportError, AssertionError): + pass + + if ( + cfg.ssd.enabled + and not self.is_mla_follower + and not (self.is_mla_backend and tp_size > 1) + and cfg.ssd.ssd_backend not in ("spdk", "spdk_proxy") + ): + rank_dir_parts = [] + if dp_rank_hint is not None: + rank_dir_parts.append(f"dp{dp_rank_hint}") + if self.pp_size > 1: + rank_dir_parts.append(f"pp{self.pp_rank}") + if self.tp_size > 1: + rank_dir_parts.append(f"tp{self.local_rank}") + if not rank_dir_parts and unique_rank != 0: + rank_dir_parts.append(f"rank{unique_rank}") + if rank_dir_parts: + cfg.ssd.storage_dir = os.path.join( + cfg.ssd.storage_dir, "_".join(rank_dir_parts) + ) + logger.info( + "UMBPStore local SSD isolation: unique_rank=%d, ssd_dir=%s", + unique_rank, + cfg.ssd.storage_dir, + ) + + if cfg.ssd.enabled and cfg.ssd.ssd_backend in ("spdk", "spdk_proxy"): + if dp_rank_hint is not None and tenant_id_base is not None: + cfg.ssd.spdk_proxy_tenant_id = tenant_id_base + dp_rank_hint + elif dp_rank_hint is not None and not explicit_tenant_id: + cfg.ssd.spdk_proxy_tenant_id = dp_rank_hint + if ( + dp_rank_hint is not None + and dp_size_hint is not None + and cfg.ssd.spdk_proxy_tenant_quota_bytes <= 0 + and dp_size_hint > 1 + ): + safe_cap = int(cfg.ssd.capacity_bytes * 0.95) + cfg.ssd.spdk_proxy_tenant_quota_bytes = max(1, safe_cap // dp_size_hint) + + self.client = UMBPClient(cfg) + if mem_pool_host is not None: + self.register_mem_pool_host(mem_pool_host) + + self.enable_pp = self.pp_size > 1 + if self.enable_pp: + self.mha_suffix = f"{self.local_rank}_{self.pp_rank}" + self.mla_suffix = f"{self.pp_rank}" + else: + self.mha_suffix = f"{self.local_rank}" + self.mla_suffix = "" + + self.split_factor = 0 + if storage_config and storage_config.should_split_heads: + self.split_factor = storage_config.tp_lcm_size // storage_config.tp_size + base_rank = self.local_rank * self.split_factor + target_ranks = [base_rank + i for i in range(self.split_factor)] + if self.enable_pp: + self.mha_suffix = [f"{rank}_{self.pp_rank}" for rank in target_ranks] + else: + self.mha_suffix = [f"{rank}" for rank in target_ranks] + + logger.info( + "UMBPStore initialized: dram=%d MB, ssd=%s, mla=%s, rank=%d, ssd_backend=%s", + cfg.dram.capacity_bytes // (1024 * 1024), + cfg.ssd.enabled, + self.is_mla_backend, + self.local_rank, + cfg.ssd.ssd_backend, + ) + + # ------------------------------------------------------------------ + # Optional KV events subscriber + # Enabled via extra_config["kv_events_subscriber"] = True/1/"true". + # Extra knobs: + # kv_events_endpoint — ZMQ connect address (default "tcp://localhost:5557") + # kv_events_topic — topic filter matching the server's topic (default "") + # ------------------------------------------------------------------ + _is_dp_mode = ( + dp_rank_hint is not None and dp_size_hint is not None and dp_size_hint > 1 + ) + _dp_rank = dp_rank_hint if dp_rank_hint is not None else 0 + + self._kv_events_subscriber: Optional[KVEventsSubscriber] = None + if _strict_bool( + extra.get("kv_events_subscriber", False), "kv_events_subscriber" + ): + # DP mode: all DP clients subscribe (filter by dp_rank in on_event). + # TP-only mode: only rank 0 subscribes. + if _is_dp_mode or self.local_rank == 0: + from sglang.srt.disaggregation.kv_events import ZmqEventPublisher + + kv_endpoint_base = str( + extra.get("kv_events_endpoint", "tcp://localhost:5557") + ) + kv_endpoint = ( + ZmqEventPublisher.offset_endpoint_port(kv_endpoint_base, _dp_rank) + if _is_dp_mode + else kv_endpoint_base + ) + kv_topic = str(extra.get("kv_events_topic", "")) + self._kv_events_subscriber = KVEventsSubscriber( + umbp_client=self.client, + endpoint=kv_endpoint, + topic=kv_topic, + dp_rank=_dp_rank if _is_dp_mode else None, + ) + self._kv_events_subscriber.start() + + # ------------------------------------------------------------------ + # Host memory pool registration + # ------------------------------------------------------------------ + def register_mem_pool_host(self, mem_pool_host: HostKVCache): + super().register_mem_pool_host(mem_pool_host) + assert self.mem_pool_host.layout in [ + "page_first", + "page_first_direct", + "page_head", + ], "UMBP store only supports page_first, page_first_direct, or page_head layout" + + # In distributed mode, pre-register the entire host KV buffer with the + # underlying RDMA IOEngine so PoolClient can take the zero-copy path + # for batch_get_into_ptr / batch_put_from_ptr (skips the staging + # buffer memcpy + lock and removes the per-call `staging_buffer_size` + # cap). Standalone returns true as no-op by IUMBPClient contract; + # we still gate on is_distributed() below to avoid a pointless call. + self._zero_copy_registered = False + if self.client is None: + return + try: + is_distributed = bool(self.client.is_distributed()) + except Exception: + is_distributed = False + if not is_distributed: + return + if not hasattr(self.client, "register_memory"): + return + if getattr(self, "_disable_zero_copy_register", False): + logger.info( + "UMBPStore: skipping host KV buffer RDMA registration because " + "disable_zero_copy_register=true (UMBP_DISABLE_ZERO_COPY_REGISTER). " + "Falling back to the staging-buffer transfer path; per-transfer " + "size is capped by distributed.staging_buffer_size." + ) + return + try: + kv_buffer = mem_pool_host.kv_buffer + host_ptr = int(kv_buffer.data_ptr()) + host_size = int(kv_buffer.numel() * kv_buffer.element_size()) + # When the buffer is backed by hugepages the mmap region is + # rounded up to the hugepage boundary. RDMA ibv_reg_mr on + # some NICs (AINIC / ROCm) requires the registered region to + # cover complete hugepages, so use the full mapped_size + # instead of the logical tensor size. + allocator = getattr(mem_pool_host, "allocator", None) + mapped_size_fn = getattr(allocator, "mapped_size_for", None) + if mapped_size_fn is not None: + mapped_size = mapped_size_fn(host_ptr) + else: + mapped_size = getattr(allocator, "mapped_size", 0) + if mapped_size > host_size: + host_size = mapped_size + ok = bool(self.client.register_memory(host_ptr, host_size)) + except Exception as exc: + logger.warning( + "UMBPStore: register_memory failed (%s); falling back to staging " + "buffer path. Per-transfer size will be capped by " + "distributed.staging_buffer_size.", + exc, + ) + return + if ok: + self._zero_copy_registered = True + logger.info( + "UMBPStore: registered host KV buffer for RDMA zero-copy " + "(ptr=0x%x, size=%d MB)", + host_ptr, + host_size // (1024 * 1024), + ) + else: + logger.warning( + "UMBPStore: register_memory returned false; staying on staging " + "buffer fallback path." + ) + + # ------------------------------------------------------------------ + # Key suffix generation — mirrors MooncakeStore + # ------------------------------------------------------------------ + def _get_mha_buffer_meta(self, keys, indices): + ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) + key_list = [] + for key_ in keys: + key_list.append(f"{key_}_{self.mha_suffix}_k") + key_list.append(f"{key_}_{self.mha_suffix}_v") + assert len(key_list) == len(ptr_list) + return key_list, ptr_list, element_size_list + + def _get_mha_split_heads_buffer_meta(self, keys, indices): + ptr_list, element_size_list = ( + self.mem_pool_host.get_split_heads_page_buffer_meta( + indices, self.split_factor + ) + ) + key_list = [] + for key_ in keys: + for suffix in self.mha_suffix: + key_list.append(f"{key_}_{suffix}_k") + key_list.append(f"{key_}_{suffix}_v") + assert len(key_list) == len(ptr_list) + return key_list, ptr_list, element_size_list + + def _get_mla_buffer_meta(self, keys, indices): + ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) + key_list = [] + for key_ in keys: + key_list.append(f"{key_}_{self.mla_suffix}_k") + assert len(key_list) == len(ptr_list) + return key_list, ptr_list, element_size_list + + def _batch_preprocess(self, keys, host_indices): + assert len(keys) > 0 + assert len(keys) == len(host_indices) // self.mem_pool_host.page_size + if self.is_mla_backend: + return self._get_mla_buffer_meta(keys, host_indices) + else: + if self.storage_config and self.storage_config.should_split_heads: + return self._get_mha_split_heads_buffer_meta(keys, host_indices) + else: + return self._get_mha_buffer_meta(keys, host_indices) + + def _batch_postprocess(self, results: List[bool], is_set_operate=False): + """Convert per-key-component results to per-page results. + + For MHA: each page has K+V → group pairs. + For MLA: each page has K only. + """ + if self.is_mla_backend: + return list(results) + else: + if self.storage_config and self.storage_config.should_split_heads: + group_size = self.split_factor * 2 + groups = [ + results[i : i + group_size] + for i in range(0, len(results), group_size) + ] + return [all(g) for g in groups] + else: + # Group K/V pairs + kv_pairs = zip(results[::2], results[1::2]) + return [k and v for k, v in kv_pairs] + + # ------------------------------------------------------------------ + # Zero-copy v1 interface + # ------------------------------------------------------------------ + def batch_get_v1( + self, + keys: List[str], + host_indices: torch.Tensor, + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> List[bool]: + key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices) + + # Normalize sizes to list of per-key sizes + if isinstance(buffer_sizes, int): + sizes = [buffer_sizes] * len(key_strs) + elif isinstance(buffer_sizes, list) and len(buffer_sizes) == 1: + sizes = buffer_sizes * len(key_strs) + else: + sizes = list(buffer_sizes) + + total_bytes = sum(sizes) + logger.debug( + "[UMBPStore] batch_get_v1: calling UMBP BatchGet: " + "keys=%d expanded_keys=%d total_bytes=%d", + len(keys), + len(key_strs), + total_bytes, + ) + get_results = self.client.batch_get_into_ptr(key_strs, list(buffer_ptrs), sizes) + success_count = sum(1 for r in get_results if r) + logger.debug( + "[UMBPStore] batch_get_v1: UMBP BatchGet done: success=%d/%d", + success_count, + len(get_results), + ) + return self._batch_postprocess(get_results) + + def _compute_expanded_depths( + self, keys: List[str], extra_info: Optional[HiCacheStorageExtraInfo] + ) -> List[int]: + """Compute per-expanded-key depth values from prefix_keys metadata. + + depth = len(prefix_keys) + page_index_within_node. + All key variants of the same page (K, V, multi-rank) share the same depth. + Returns an empty list if no metadata is available (caller falls back to plain LRU). + """ + prefix_keys = getattr(extra_info, "prefix_keys", None) if extra_info else None + if prefix_keys is None: + return [] + + prefix_len = len(prefix_keys) + depths_per_page = [prefix_len + i for i in range(len(keys))] + + # Expand to match the key_strs layout produced by _batch_preprocess. + expanded = [] + for d in depths_per_page: + if self.is_mla_backend: + expanded.append(d) # MLA: 1 key per page + elif self.storage_config and self.storage_config.should_split_heads: + # split heads: 2 keys per split rank, split_factor ranks per page + for _ in range(self.split_factor): + expanded.append(d) + expanded.append(d) + else: + expanded.append(d) # K + expanded.append(d) # V + return expanded + + def batch_set_v1( + self, + keys: List[str], + host_indices: torch.Tensor, + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> List[bool]: + # Follower never writes (CacheController also sets backup_skip, but guard here too) + if self.is_mla_follower: + page_count = len(host_indices) // self.mem_pool_host.page_size + return [True] * page_count + + key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices) + + if isinstance(buffer_sizes, int): + sizes = [buffer_sizes] * len(key_strs) + elif isinstance(buffer_sizes, list) and len(buffer_sizes) == 1: + sizes = buffer_sizes * len(key_strs) + else: + sizes = list(buffer_sizes) + + expanded_depths = self._compute_expanded_depths(keys, extra_info) + + total_bytes = sum(sizes) + logger.debug( + "[UMBPStore] batch_set_v1: calling UMBP BatchPut: " + "keys=%d expanded_keys=%d total_bytes=%d with_depth=%s", + len(keys), + len(key_strs), + total_bytes, + bool(expanded_depths), + ) + + if expanded_depths: + put_results = self.client.batch_put_from_ptr_with_depth( + key_strs, list(buffer_ptrs), sizes, expanded_depths + ) + else: + put_results = self.client.batch_put_from_ptr( + key_strs, list(buffer_ptrs), sizes + ) + + success_count = sum(1 for r in put_results if r) + logger.debug( + "[UMBPStore] batch_set_v1: UMBP BatchPut done: success=%d/%d", + success_count, + len(put_results), + ) + return self._batch_postprocess(put_results, is_set_operate=True) + + def batch_exists( + self, keys: List[str], extra_info: Optional[HiCacheStorageExtraInfo] = None + ) -> int: + """Return count of consecutive existing keys from start.""" + if self.is_mla_backend: + query_keys = [f"{key}_{self.mla_suffix}_k" for key in keys] + key_multiplier = 1 + else: + query_keys = [] + if self.storage_config and self.storage_config.should_split_heads: + for key in keys: + for suffix in self.mha_suffix: + query_keys.append(f"{key}_{suffix}_k") + query_keys.append(f"{key}_{suffix}_v") + key_multiplier = 2 * self.split_factor + else: + for key in keys: + query_keys.append(f"{key}_{self.mha_suffix}_k") + query_keys.append(f"{key}_{self.mha_suffix}_v") + key_multiplier = 2 + + hit_count = self.client.batch_exists_consecutive(query_keys) + return hit_count // key_multiplier + + # ------------------------------------------------------------------ + # Legacy ABC interface (required by HiCacheStorage) + # ------------------------------------------------------------------ + def get( + self, + key: str, + target_location: Optional[Any] = None, + target_sizes: Optional[Any] = None, + ) -> torch.Tensor | None: + if target_location is None or target_sizes is None: + return None + ok = self.client.get_into_ptr(key, target_location, target_sizes) + return target_location if ok else None + + def batch_get( + self, + keys: List[str], + target_locations: Optional[Any] = None, + target_sizes: Optional[Any] = None, + ) -> int: + if not keys: + return 0 + assert len(keys) == len(target_locations) == len(target_sizes) + results = self.client.batch_get_into_ptr( + keys, + list(target_locations), + list(target_sizes), + ) + for i, ok in enumerate(results): + if not ok: + return i + return len(keys) + + def set( + self, + key: str, + value: Optional[Any] = None, + target_location: Optional[Any] = None, + target_sizes: Optional[Any] = None, + ) -> bool: + if self.is_mla_follower: + return True + if target_location is None or target_sizes is None: + return False + return self.client.put_from_ptr(key, target_location, target_sizes) + + def batch_set( + self, + keys: List[str], + values: Optional[Any] = None, + target_locations: Optional[Any] = None, + target_sizes: Optional[Any] = None, + ) -> bool: + if not keys: + return False + if self.is_mla_follower: + return True + assert len(keys) == len(target_locations) == len(target_sizes) + results = self.client.batch_put_from_ptr( + keys, + list(target_locations), + list(target_sizes), + ) + return all(results) + + def exists(self, key: str) -> bool: + return self.client.exists(key) + + def clear(self) -> None: + self.client.clear() + + def flush(self) -> bool: + if self.client is None or not hasattr(self.client, "flush"): + return True + return bool(self.client.flush()) + + def close(self) -> None: + if getattr(self, "_kv_events_subscriber", None) is not None: + try: + self._kv_events_subscriber.stop() + except Exception: + logger.exception("KVEventsSubscriber stop during close failed") + self._kv_events_subscriber = None + if getattr(self, "client", None) is None: + return + try: + self.flush() + except Exception: + logger.exception("UMBPStore flush during close failed") + self.client = None + + +class KVEventsSubscriber: + """Subscribe to SGLang KV cache events published over ZMQ and forward + them to the UMBP Master via ``umbp_client``. + + Runs a background thread that receives :class:`KVEventBatch` messages + from a ``ZmqEventPublisher`` and dispatches each individual event to + :meth:`on_event`. + + Parameters + ---------- + umbp_client: + The ``UMBPClient`` instance owned by the parent ``UMBPStore``. + Will be used to publish events to the UMBP Master. + endpoint: + ZMQ endpoint of the SGLang publisher, e.g. ``"tcp://localhost:5557"``. + For DP attention rank *N*, the port is offset by *N* (rank 0 → 5557, + rank 1 → 5558, …). + topic: + Topic filter passed to ``zmq.SUBSCRIBE``. Must match the ``topic`` + set in ``KVEventsConfig`` on the server side. Empty string subscribes + to everything. + poll_timeout_ms: + How long (in milliseconds) the receiver thread waits for a message + before looping and checking the stop flag. + """ + + def __init__( + self, + umbp_client: Any, + endpoint: str = "tcp://localhost:5557", + topic: str = "", + poll_timeout_ms: int = 100, + dp_rank: Optional[int] = None, + ) -> None: + self._umbp_client = umbp_client + self._endpoint = endpoint + self._topic = topic + self._poll_timeout_ms = poll_timeout_ms + # When set, only process events whose attn_dp_rank matches (DP mode). + self._dp_rank = dp_rank + self._stop_event = threading.Event() + self._thread: Optional[threading.Thread] = None + # Cache UMBPTierType constants to avoid repeated attribute lookups. + import mori.umbp as _umbp_mod + + _tier = _umbp_mod.UMBPTierType + self._tier_hbm = _tier.HBM + self._tier_dram = _tier.DRAM + + # ------------------------------------------------------------------ + # Lifecycle + # ------------------------------------------------------------------ + + def start(self) -> None: + """Start the subscriber background thread.""" + if self._thread is not None and self._thread.is_alive(): + return + self._stop_event.clear() + self._thread = threading.Thread( + target=self._run, + daemon=True, + name="kv-events-subscriber", + ) + self._thread.start() + logger.info( + "KVEventsSubscriber started: endpoint=%s, topic=%r", + self._endpoint, + self._topic, + ) + + def stop(self, timeout: float = 2.0) -> None: + """Signal the subscriber thread to stop and wait for it to finish.""" + self._stop_event.set() + if self._thread is not None: + self._thread.join(timeout=timeout) + self._thread = None + logger.info("KVEventsSubscriber stopped") + + # ------------------------------------------------------------------ + # Event handler + # ------------------------------------------------------------------ + + def _medium_to_tier(self, medium: Optional[str]) -> Any: + """Map a KV event medium string to a UMBPTierType value. + + ``StorageMedium.GPU`` ("GPU") maps to HBM (GPU on-chip memory). + All other values (CPU pinned, unknown) map to DRAM. + """ + # kv_events.py exposes the enum class ``StorageMedium`` (not a + # module-level ``MEDIUM_GPU`` constant — that import was removed when + # the file was reorganised but this consumer was never updated, so + # PD-disagg paths crash here when KVEventsSubscriber receives the + # first GPU-tier event). Compare against the enum value's string. + from sglang.srt.disaggregation.kv_events import StorageMedium + + return self._tier_hbm if medium == StorageMedium.GPU.value else self._tier_dram + + def on_event( + self, event: Any, batch_ts: float, attn_dp_rank: Optional[int] + ) -> None: + """Called once per individual KV cache event. + + Translates SGLang KV cache events into UMBP external KV block + report / revoke calls so the UMBP Master can track which GPU/CPU + blocks are available on each node for cross-node cache lookup. + + Parameters + ---------- + event: + One of :class:`~sglang.srt.disaggregation.kv_events.BlockStored`, + :class:`~sglang.srt.disaggregation.kv_events.BlockRemoved`, or + :class:`~sglang.srt.disaggregation.kv_events.AllBlocksCleared`. + batch_ts: + Unix timestamp of the :class:`KVEventBatch` that contained this event. + attn_dp_rank: + DP attention rank that produced the event, or ``None`` for rank 0. + """ + from sglang.srt.disaggregation.kv_events import ( + AllBlocksCleared, + BlockRemoved, + BlockStored, + ) + + if self._dp_rank is not None: + event_dp_rank = attn_dp_rank if attn_dp_rank is not None else 0 + if event_dp_rank != self._dp_rank: + return + + if isinstance(event, BlockStored): + hashes = [str(h) for h in event.block_hashes] + tier = self._medium_to_tier(event.medium) + ok = self._umbp_client.report_external_kv_blocks(hashes, tier) + if not ok: + logger.warning( + "report_external_kv_blocks failed for %d hashes (dp_rank=%s)", + len(hashes), + attn_dp_rank, + ) + + elif isinstance(event, BlockRemoved): + hashes = [str(h) for h in event.block_hashes] + tier = self._medium_to_tier(event.medium) + ok = self._umbp_client.revoke_external_kv_blocks(hashes, tier) + if not ok: + logger.warning( + "revoke_external_kv_blocks failed for %d hashes (dp_rank=%s)", + len(hashes), + attn_dp_rank, + ) + + elif isinstance(event, AllBlocksCleared): + # AllBlocksCleared wipes the local KV event publisher's cache. Ask + # the master to clear this node's whole bucket for each tier this + # subscriber can report, instead of sending a full hash list back. + ok_hbm = self._umbp_client.revoke_all_external_kv_blocks_at_tier( + self._tier_hbm + ) + ok_dram = self._umbp_client.revoke_all_external_kv_blocks_at_tier( + self._tier_dram + ) + if ok_hbm is False or ok_dram is False: + logger.warning( + "revoke_all_external_kv_blocks_at_tier (all-clear) failed " + "(hbm=%s dram=%s)", + ok_hbm, + ok_dram, + ) + + else: + logger.debug( + "KVEventsSubscriber: unhandled event type %s", type(event).__name__ + ) + + # ------------------------------------------------------------------ + # Internal + # ------------------------------------------------------------------ + + def _run(self) -> None: + import zmq + from msgspec.msgpack import Decoder + + from sglang.srt.disaggregation.kv_events import KVEventBatch + + decoder = Decoder(type=KVEventBatch) + ctx = zmq.Context.instance() + sub = ctx.socket(zmq.SUB) + sub.connect(self._endpoint) + sub.setsockopt_string(zmq.SUBSCRIBE, self._topic) + logger.debug( + "KVEventsSubscriber connected to %s, topic=%r", self._endpoint, self._topic + ) + + try: + while not self._stop_event.is_set(): + if not sub.poll(self._poll_timeout_ms): + continue + try: + parts = sub.recv_multipart() + # Publisher sends: [topic, seq_bytes, payload] + if len(parts) != 3: + logger.warning( + "Unexpected frame count %d from publisher", len(parts) + ) + continue + _, _seq_bytes, payload = parts + batch = decoder.decode(payload) + for event in batch.events: + self.on_event(event, batch.ts, batch.attn_dp_rank) + except Exception: + logger.exception("KVEventsSubscriber error decoding message") + finally: + sub.close(linger=0) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 4b58e69d4..7952c1e63 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1923,6 +1923,7 @@ class ServerArgs: "dynamic", "eic", "simm", + "mori", ], ), ] = None diff --git a/test/registered/unit/mem_cache/test_umbp_host_allocator.py b/test/registered/unit/mem_cache/test_umbp_host_allocator.py new file mode 100644 index 000000000..4e64b5a97 --- /dev/null +++ b/test/registered/unit/mem_cache/test_umbp_host_allocator.py @@ -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() diff --git a/test/registered/unit/mem_cache/test_umbp_store.py b/test/registered/unit/mem_cache/test_umbp_store.py new file mode 100755 index 000000000..142840a76 --- /dev/null +++ b/test/registered/unit/mem_cache/test_umbp_store.py @@ -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 -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()