From f1f2380d2b425563980ef1d1d0590ef0daa7bf04 Mon Sep 17 00:00:00 2001 From: Niko Ma Date: Sat, 5 Sep 2026 09:21:43 +0800 Subject: [PATCH] [Unified Cache][6/N]: Add UMBP external linker (#37578) Co-authored-by: Zhangheng --- python/sglang/srt/mem_cache/registry.py | 6 + .../storage/umbp/umbp_direct_linker.py | 1050 +++++++++++++++++ .../storage/umbp/umbp_host_allocator.py | 40 +- .../srt/mem_cache/storage/umbp/umbp_store.py | 550 ++++++--- python/sglang/srt/server_args.py | 2 +- .../test_hicache_storage_umbp_backend.py | 1 + .../unit/mem_cache/test_registry.py | 51 + .../mem_cache/test_umbp_host_allocator.py | 38 +- .../unit/mem_cache/test_umbp_store.py | 53 + 9 files changed, 1584 insertions(+), 207 deletions(-) create mode 100644 python/sglang/srt/mem_cache/storage/umbp/umbp_direct_linker.py diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index 8e90b901d..2e6c0513f 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -204,6 +204,12 @@ def _create_unified_radix_cache( ) linker_cls = MooncakeDirectLinker + elif backend == "mori": + from sglang.srt.mem_cache.storage.umbp.umbp_direct_linker import ( + UMBPDirectLinker, + ) + + linker_cls = UMBPDirectLinker else: raise ValueError( f"Unknown unified cache external linker backend: {backend!r}" diff --git a/python/sglang/srt/mem_cache/storage/umbp/umbp_direct_linker.py b/python/sglang/srt/mem_cache/storage/umbp/umbp_direct_linker.py new file mode 100644 index 000000000..39c70a5d1 --- /dev/null +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_direct_linker.py @@ -0,0 +1,1050 @@ +from __future__ import annotations + +import logging +import os +import threading +from collections import defaultdict +from concurrent.futures import Future +from dataclasses import dataclass +from queue import Empty, Queue +from typing import Any, Callable + +import numpy as np +import torch + +from sglang.srt.mem_cache.cache_init_params import CacheInitParams +from sglang.srt.mem_cache.hicache_storage import ( + HiCacheStorageConfig, + PoolHitPolicy, + PoolName, + PoolTransfer, +) +from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import ( + resolve_hybrid_device_pool_group, +) +from sglang.srt.mem_cache.unified_cache.unified_cache_linker import UnifiedCacheLinker +from sglang.srt.runtime_context import get_memory, get_model +from sglang.srt.utils import freeze_gc, get_device_module + +logger = logging.getLogger(__name__) +device_module = get_device_module() + +# Keep every control-plane RPC comfortably below gRPC's message-size limit. +# This is a logical-page count; existence queries carry keys only, no ranges. +CHUNK_PAGES = 64 + +# Budget by ranges because pool layouts attach different counts per object; +# 8192 stays below gRPC's default message limit. +RANGES_PER_CALL = int(os.getenv("UMBP_RANGES_PER_CALL", "8192")) + + +def _ordered_layers(entry) -> list[int]: + component_lengths = {len(component) for component in entry.components} + if len(component_lengths) != 1: + raise ValueError( + f"UMBP pool {entry.name} components have different layer counts." + ) + pool_layer_count = component_lengths.pop() + if pool_layer_count != len(entry.layer_mapping): + raise ValueError( + f"UMBP pool {entry.name} has {pool_layer_count} buffers per component " + f"but {len(entry.layer_mapping)} mapped layers." + ) + by_buffer = { + buffer_index: logical_layer + for logical_layer, buffer_index in entry.layer_mapping.items() + } + if sorted(by_buffer) != list(range(pool_layer_count)): + raise ValueError( + f"UMBP pool {entry.name} layer mapping is not a contiguous bijection." + ) + return [by_buffer[index] for index in range(pool_layer_count)] + + +class LayerWiseLoadCounter: + """CPU completion counter compatible with KV pools' layer wait hook.""" + + def __init__(self, num_layers: int): + self.num_layers = num_layers + self._producer_index = -1 + self.consumer_index = -1 + self._futures: dict[int, list[Future]] = {} + + def update_producer(self) -> int: + self._producer_index += 1 + self._futures[self._producer_index] = [Future() for _ in range(self.num_layers)] + return self._producer_index + + def set_consumer(self, index: int) -> None: + self.consumer_index = index + + def complete(self, index: int, layer: int) -> None: + self._futures[index][layer].set_result(None) + + def fail(self, index: int, error: BaseException) -> None: + for future in self._futures.get(index, ()): + if not future.done(): + future.set_exception(error) + + def wait_until(self, threshold: int) -> None: + index = self.consumer_index + futures = self._futures.get(index) + if futures is None: + return + try: + futures[threshold].result() + except BaseException as error: + raise RuntimeError("UMBP layer-wise KV load failed.") from error + finally: + if threshold == self.num_layers - 1: + self._futures.pop(index, None) + + def reset(self) -> None: + self._producer_index = -1 + self.consumer_index = -1 + self._futures.clear() + + +@dataclass +class _PoolRangePlan: + """Object keys and locations for one pool load.""" + + name: PoolName + keys: list[str] + locations: list[int] + entries_per_page: int + + +# One queued offload: the pools it resolved to, and the event guarding its KV. +_OffloadTask = tuple[list[PoolTransfer], object] + + +def _offload_task_pages(expanded: list[PoolTransfer]) -> int: + """Pages this task puts into its widest pool, which is what sizes a plan.""" + return max((len(transfer.keys or ()) for transfer in expanded), default=0) + + +def _object_sizes_per_page(entry) -> list[int]: + """Return per-page object sizes independently of emitted ranges. + + This keeps the tier's exact-tiling validation independent of range generation. + """ + if entry.packed: + return [ + sum(size for component in entry.buffer_meta for _, _, size in component) + ] + return [sum(size for _, _, size in component) for component in entry.buffer_meta] + + +def _config_bool(value: Any, key: str) -> bool: + if isinstance(value, bool): + return value + if isinstance(value, int) and value in (0, 1): + return bool(value) + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + raise ValueError(f"UMBP linker config {key!r} must be boolean, got {value!r}.") + + +def _materialize_cpu_indices(indices: torch.Tensor) -> torch.Tensor: + """Materialize CPU indices used to derive per-pool row locations.""" + return indices.detach().to(device="cpu", dtype=torch.int64).flatten() + + +def _parse_storage_extra_config(raw_config): + # Keep the linker module importable in CPU-only unit tests. The hybrid + # controller imports device-specific memory-pool modules transitively. + from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( + HybridCacheController, + ) + + extra_config, *_ = HybridCacheController.parse_storage_backend_extra_config( + raw_config + ) + return extra_config + + +class UMBPDirectLinker(UnifiedCacheLinker): + def __init__( + self, + server_args, + params: CacheInitParams, + *, + components, + _storage=None, + ): + self.page_size = params.page_size + # Group layers to amortize per-object RPC overhead; 8 is the measured default. + self.layer_group = max(1, int(os.getenv("UMBP_LAYER_GROUP", "8"))) + # Coalesce queued offload tasks up to this many pages; offload_nodes + # queues one task per node, so they arrive a page or two at a time. + self._offload_coalesce_pages = max( + 1, int(os.getenv("UMBP_OFFLOAD_COALESCE_PAGES", "1024")) + ) + if _config_bool(os.getenv("UMBP_LOAD_SPLIT") or "0", "UMBP_LOAD_SPLIT"): + raise ValueError( + "UMBP_LOAD_SPLIT is not supported by the dedup-after-insert " + "load flow: ranks may receive different page sets and deadlock." + ) + + kvcache = params.token_to_kv_pool_allocator.get_kvcache() + self._async_offload_index_snapshot = True + self._offload_index_fallback_warned = False + self._offload_index_stream = None + self._offload_index_done = None + self._offload_index_device = None + self._offload_index_buffers: list[torch.Tensor | None] = [] + distributed = ( + torch.distributed.is_available() and torch.distributed.is_initialized() + ) + tp_rank = 0 + if distributed: + tp_rank = torch.distributed.get_rank(group=params.tp_cache_group) + self.pool_group = resolve_hybrid_device_pool_group( + kvcache=kvcache, + page_size=self.page_size, + params=params, + components=components, + ) + self.pools = self.pool_group.entry_map + self.num_layers = self.pool_group.num_layers + if self.num_layers <= 0: + raise ValueError("UMBP requires at least one logical layer.") + self.pool_layers = { + name: _ordered_layers(entry) for name, entry in self.pools.items() + } + invalid_layers = { + name: [layer for layer in layers if not 0 <= layer < self.num_layers] + for name, layers in self.pool_layers.items() + } + invalid_layers = { + name: layers for name, layers in invalid_layers.items() if layers + } + if invalid_layers: + raise ValueError( + f"UMBP pool mappings contain out-of-range logical layers: {invalid_layers}." + ) + extra_config = _parse_storage_extra_config( + get_memory().hicache_storage_backend_extra_config + ) + extra_config = dict(extra_config) + standalone_requested = bool( + extra_config.get("standalone_address") + or os.getenv("UMBP_STANDALONE_ADDRESS") + ) + if "ssd_enabled" in extra_config and _config_bool( + extra_config["ssd_enabled"], "ssd_enabled" + ): + raise ValueError( + "Direct UMBP requires ssd_enabled=false because its GPU path " + "cannot use the corresponding host-memory fallback." + ) + extra_config["ssd_enabled"] = False + + if "cache_remote_fetches" in extra_config and _config_bool( + extra_config["cache_remote_fetches"], "cache_remote_fetches" + ): + raise ValueError( + "Direct UMBP requires cache_remote_fetches=false because its GPU " + "path cannot use the corresponding host-memory fallback." + ) + if standalone_requested: + extra_config.pop("cache_remote_fetches", None) + else: + extra_config["cache_remote_fetches"] = False + + min_object_size = min( + min( + pool.get_page_buffer_meta( + torch.arange(pool.page_size, dtype=torch.int64) + )[1] + ) + for pool in self.pools.values() + ) + if standalone_requested: + extra_config.pop("dram_page_size", None) + else: + dram_page_size = int(extra_config.get("dram_page_size", min_object_size)) + if not 0 < dram_page_size <= min_object_size: + raise ValueError( + "Direct UMBP requires 0 < dram_page_size <= the smallest " + f"per-layer object ({min_object_size} bytes), got {dram_page_size}." + ) + extra_config["dram_page_size"] = dram_page_size + + storage_config = HiCacheStorageConfig( + tp_rank=tp_rank, + tp_size=server_args.tp_size, + pp_rank=params.pp_rank, + pp_size=params.pp_size, + attn_cp_rank=params.attn_cp_rank, + attn_cp_size=params.attn_cp_size, + is_mla_model=True, + enable_storage_metrics=False, + is_page_first_layout=False, + model_name=get_model().model_path, + extra_config=extra_config, + ) + + if _storage is None: + from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore + + # per_rank_keyspace: every object key this class writes carries a + # tp{rank} suffix (set just below), so the store must not put the + # ranks into the shared-SSD leader/follower scheme meant for + # deduplicating replicated MLA KV. + self.storage = UMBPStore( + storage_config, mem_pool_host=None, per_rank_keyspace=True + ) + else: + self.storage = _storage + + try: + client = self.storage.client + mode = client.get_deployment_mode() + mode_type = type(mode) + # What page-granular objects actually need is ranged multi-buffer + # I/O. That used to be true only of StandaloneProcess, so the gate + # was written as a mode test -- but mori now implements it in the + # in-process client too ("Both media behind LocalStorageManager + # implement ranged I/O now", standalone_client.h), and a mode test + # would keep rejecting a client that can do the job. + # + # So ask the client what it supports instead of inferring it from + # which mode it is. supports_ranged_io() is the capability this + # code depends on, it is already consulted below, and a client that + # answers truthfully cannot be wrongly admitted by it. + supports_ranged = getattr(client, "supports_ranged_io", None) + if not callable(supports_ranged) or not bool(supports_ranged()): + raise ValueError( + f"Direct UMBP needs ranged multi-buffer I/O, which this " + f"{mode!r} client does not advertise. Upgrade mori; if the " + "server has a Distributed inner backend, set " + "UMBP_DISTRIBUTED_RANGED_SCRATCH_BYTES to a positive value " + "(both scratch arenas must be non-zero), and note that a " + "read-only SharedSSDFollower is whole-object by design." + ) + get_backend_mode = getattr(client, "get_backend_mode", None) + self.backend_mode = ( + get_backend_mode() if callable(get_backend_mode) else None + ) + self.deployment_mode = mode + self._standalone_process_mode = mode == mode_type.StandaloneProcess + if getattr(self.storage, "_disable_zero_copy_register", False): + raise ValueError( + "Direct UMBP cannot disable zero-copy memory registration." + ) + + self.storage.mem_pool_host = self.pool_group + self.storage._kv_anchor_is_logical = True + self.storage.registered_pools = self.pools + rank_suffix = f"tp{tp_rank}_cp{params.attn_cp_rank}_pp{params.pp_rank}" + self.storage.mla_suffix = rank_suffix + self.storage.mha_suffix = rank_suffix + self._register_buffers() + # Report the mode rather than asserting it in the text: the line is + # what every acceptance check greps to prove the linker attached at + # all, and it used to say "standalone_process" unconditionally, so + # it could not have shown an embedded run for what it was. + logger.info( + "UMBPDirectLinker topology=%s+%s ranged_io=yes", + mode.name, + self.backend_mode.name if self.backend_mode is not None else None, + ) + except BaseException: + self.storage.close() + raise + + self.layer_done_counter = LayerWiseLoadCounter(self.num_layers) + if PoolName.MAMBA in self.pools: + params.req_to_token_pool.register_layer_transfer_counter( + self.layer_done_counter + ) + self._pending: dict[str, list[PoolTransfer]] = {} + self._gc_frozen = False + self._load_queue: Queue[ + tuple[int, list[str], list[_PoolRangePlan], object] | None + ] = Queue() + self._completed_loads: Queue[list[str]] = Queue() + self._offload_queue: Queue[tuple[list[PoolTransfer], object] | None] = Queue() + self._offload_results: Queue[bool] = Queue() + self._stats = { + "lookup": 0, + "load": 0, + "offload": 0, + # offload / offload_batches is the coalescing actually achieved. + "offload_batches": 0, + } + self._load_thread = threading.Thread( + target=self._load_thread_func, + daemon=True, + name=f"umbp-load-tp{tp_rank}", + ) + self._offload_thread = threading.Thread( + target=self._offload_thread_func, + daemon=True, + name=f"umbp-offload-tp{tp_rank}", + ) + self._closed = False + self._load_thread.start() + self._offload_thread.start() + + def _register_buffers(self) -> None: + seen = set() + self._registered: list[tuple[int, int]] = [] + for pool in self.pools.values(): + for buffer in pool.get_hybrid_pool_buffer(): + storage = buffer.untyped_storage() + allocation = (int(storage.data_ptr()), int(storage.nbytes())) + if allocation in seen: + continue + seen.add(allocation) + if not self.storage.client.register_memory(*allocation): + raise RuntimeError( + "Failed to register a GPU KV buffer with UMBP: " + f"ptr=0x{allocation[0]:x}, size={allocation[1]}." + ) + self._registered.append(allocation) + + def _object_keys_for_pages( + self, page_keys: list[str], transfer: PoolTransfer + ) -> tuple[list[str], int]: + component_keys, multiplier = self.storage._get_hybrid_page_component_keys( + page_keys, transfer + ) + entry = self.pools[transfer.name] + # One key names one stored object, and a packed entry stores a whole + # page as one object regardless of how many components it has. + entries_per_page = 1 if entry.packed else len(entry.components) + if multiplier != entries_per_page: + raise ValueError( + f"UMBP pool {transfer.name} produced {multiplier} keys per page " + f"but its layout yields {entries_per_page} objects per page " + f"(packed={entry.packed}, components={len(entry.components)})." + ) + # No layer suffix: one object per page (or per page component) holds + # every layer, and a layer is read back as a byte range inside it. + return component_keys, multiplier + + def _object_keys(self, transfer: PoolTransfer) -> list[str]: + keys, _ = self._object_keys_for_pages(list(transfer.keys or []), transfer) + return keys + + def _page_exists(self, page_keys: list[str], transfer: PoolTransfer) -> list[bool]: + entry = self.pools[transfer.name] + objects_per_page = 1 if entry.packed else len(entry.components) + max_objects = CHUNK_PAGES * self.num_layers + pages_per_call = max(1, max_objects // objects_per_page) + + page_exists = [] + for start in range(0, len(page_keys), pages_per_call): + chunk_pages = page_keys[start : start + pages_per_call] + object_keys, _ = self._object_keys_for_pages(chunk_pages, transfer) + exists = list(self.storage.client.batch_exists(object_keys)) + if len(exists) != len(object_keys): + raise RuntimeError( + f"UMBP exists result-size mismatch for pool {transfer.name}: " + f"expected={len(object_keys)} actual={len(exists)}." + ) + page_exists.extend( + all(exists[index : index + objects_per_page]) + for index in range(0, len(exists), objects_per_page) + ) + return page_exists + + @staticmethod + def _apply_hit_policy( + valid_pages: list[int], page_exists: list[bool], transfer: PoolTransfer + ) -> list[int]: + present_prefix = [0] + for present in page_exists: + present_prefix.append(present_prefix[-1] + int(present)) + + if transfer.hit_policy == PoolHitPolicy.ALL_PAGES: + return [end for end in valid_pages if present_prefix[end] == end] + if transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES: + trailing = max(1, len(transfer.keys or ())) + return [ + end + for end in valid_pages + if present_prefix[end] - present_prefix[max(0, end - trailing)] + == end - max(0, end - trailing) + ] + raise ValueError(f"Unsupported pool hit policy: {transfer.hit_policy}") + + def lookup(self, rid: str, transfers: list[PoolTransfer]) -> list[int]: + expanded = self.pool_group.resolve_transfers(transfers) + if not expanded: + return [] + kv = next(transfer for transfer in transfers if transfer.name == PoolName.KV) + page_keys = list(kv.keys or []) + if not page_keys: + return [] + + valid_pages = list(range(1, len(page_keys) + 1)) + for transfer in expanded: + # Probe only as far as the surviving boundary: no hit policy reads + # past its own end offset, so pages beyond the longest candidate + # cannot change the answer. This runs synchronously inside the + # scheduler's prefill batch build, and DP-attention ranks are + # lockstep, so every extra key is stall charged to all of them. + page_exists = self._page_exists(page_keys[: valid_pages[-1]], transfer) + valid_pages = self._apply_hit_policy(valid_pages, page_exists, transfer) + if not valid_pages: + break + + self._stats["lookup"] += 1 + if valid_pages: + logger.debug( + "UMBP direct linker lookup hit: rid=%s pages=%d candidates=%d", + rid, + valid_pages[-1], + len(valid_pages), + ) + return valid_pages + + def load(self, rid: str, transfers: list[PoolTransfer]) -> bool: + # Lookup establishes a restorable boundary before insert de-duplicates + # resident pages. The remaining transfer can therefore contain only a + # side pool such as SWA, with no KV transfer at all. + expanded = self.pool_group.resolve_transfers( + transfers, allow_partial=True, allow_missing_kv=True + ) + if not expanded: + return False + if rid in self._pending: + raise RuntimeError(f"UMBP load for rid={rid} is already queued.") + self._pending[rid] = expanded + return True + + def cancel_queued_load(self, rid: str) -> bool: + # The tree node is already visible in L1. Dropping its transfer would + # leave a device hit pointing at slots that were never populated. + return False + + def num_completed_loads(self) -> int: + return self._completed_loads.qsize() + + def pop_completed_load(self) -> list[str]: + return self._completed_loads.get_nowait() + + def start_layer_wise_loading(self) -> int: + if not self._pending: + return -1 + self._freeze_gc_once() + pending = self._pending + rids = list(pending) + plans = self._build_load_plans(list(pending.values())) + ready_event = device_module.Event() + ready_event.record() + counter_index = self.layer_done_counter.update_producer() + self._load_queue.put((counter_index, rids, plans, ready_event)) + self._pending = {} + self._stats["load"] += len(pending) + return counter_index + + def _build_load_plans( + self, + request_transfers: list[list[PoolTransfer]], + *, + materialize_indices: Callable[[torch.Tensor], torch.Tensor] | None = None, + ) -> list[_PoolRangePlan]: + """Build a batch plan shared by load and offload.""" + grouped: dict[PoolName, list[PoolTransfer]] = {} + for transfers in request_transfers: + for transfer in transfers: + grouped.setdefault(transfer.name, []).append(transfer) + + plans = [] + # One logical source can fan out to several physical pools (for + # example GLM's KV and INDEXER pools). Snapshot it once, then let each + # pool independently validate and derive its own row geometry. + cpu_indices: dict[int, torch.Tensor] = {} + for name, transfers in grouped.items(): + entry = self.pools[name] + entries_per_page = 1 if entry.packed else len(entry.components) + keys: list[str] = [] + locations: list[int] = [] + for transfer in transfers: + page_keys = list(transfer.keys or []) + transfer_keys, multiplier = self._object_keys_for_pages( + page_keys, transfer + ) + if multiplier != entries_per_page: + raise ValueError( + f"UMBP pool {name} emits {multiplier} keys per page but " + f"its layout yields {entries_per_page} objects per page " + f"(packed={entry.packed})." + ) + if len(transfer_keys) != len(page_keys) * entries_per_page: + raise ValueError( + f"UMBP pool {name} key count mismatch: " + f"keys={len(transfer_keys)} pages={len(page_keys)}." + ) + keys.extend(transfer_keys) + indices = transfer.host_indices + if indices is None: + raise ValueError(f"UMBP pool {name} transfer has no indices.") + source_id = id(indices) + prepared_indices = cpu_indices.get(source_id) + if prepared_indices is None: + prepared_indices = ( + materialize_indices(indices) + if materialize_indices is not None + else _materialize_cpu_indices(indices) + ) + cpu_indices[source_id] = prepared_indices + locations.extend(entry.prepare_locations(prepared_indices)) + if len(keys) != len(locations) * entries_per_page: + raise ValueError( + f"UMBP pool {name} plan mismatch: keys={len(keys)} " + f"rows={len(locations)} per_page={entries_per_page}." + ) + plans.append(_PoolRangePlan(name, keys, locations, entries_per_page)) + + if not plans or not plans[0].keys: + raise ValueError("Layer-wise UMBP load has no object keys.") + return plans + + def _materialize_offload_indices( + self, indices: torch.Tensor, slot: int + ) -> torch.Tensor: + if not self._async_offload_index_snapshot or indices.device.type == "cpu": + return _materialize_cpu_indices(indices) + if ( + indices.dtype != torch.int64 + or indices.ndim != 1 + or not indices.is_contiguous() + ): + return self._fallback_offload_indices(indices, "tensor shape or dtype") + if ( + self._offload_index_device is not None + and indices.device != self._offload_index_device + ): + return self._fallback_offload_indices(indices, "device mismatch") + + try: + with device_module.device(indices.device): + if self._offload_index_stream is None: + self._offload_index_device = indices.device + self._offload_index_stream = device_module.Stream( + device=indices.device + ) + self._offload_index_done = device_module.Event() + while len(self._offload_index_buffers) <= slot: + self._offload_index_buffers.append(None) + count = indices.numel() + buffer = self._offload_index_buffers[slot] + if buffer is None or buffer.numel() < count: + buffer = self._allocate_offload_index_buffer(count) + self._offload_index_buffers[slot] = buffer + source = indices.detach() + with device_module.stream(self._offload_index_stream): + # Register before enqueue so a failed copy cannot leave an + # untracked side-stream read of the source allocation. + source.record_stream(self._offload_index_stream) + buffer[:count].copy_(source, non_blocking=True) + self._offload_index_done.record(self._offload_index_stream) + self._offload_index_done.synchronize() + return buffer[:count] + except RuntimeError: + self._async_offload_index_snapshot = False + logger.exception( + "UMBP async index snapshot failed; falling back to synchronous D2H" + ) + return _materialize_cpu_indices(indices) + + @staticmethod + def _allocate_offload_index_buffer(count: int) -> torch.Tensor: + return torch.empty(count, dtype=torch.int64, device="cpu", pin_memory=True) + + def _fallback_offload_indices( + self, indices: torch.Tensor, reason: str + ) -> torch.Tensor: + if not self._offload_index_fallback_warned: + self._offload_index_fallback_warned = True + logger.warning( + "UMBP offload index snapshot is using synchronous fallback: " + "%s (device=%s dtype=%s shape=%s contiguous=%s)", + reason, + indices.device, + indices.dtype, + tuple(indices.shape), + indices.is_contiguous(), + ) + return _materialize_cpu_indices(indices) + + def _load_thread_func(self) -> None: + while True: + task = self._load_queue.get() + try: + if task is None: + return + counter_index, rids, plans, ready_event = task + try: + self._run_layer_wise_batch(counter_index, plans, ready_event) + finally: + self._completed_loads.put(rids) + finally: + self._load_queue.task_done() + + def _all_layer_ranges(self, plan: _PoolRangePlan): + """Every layer's ranges, accumulated per object. + + Offload requires one call to carry ranges that tile the object exactly, + so an object's ranges must never be split across calls. + + The group is the pool's own layer stack, so it always covers the pool + and the None return is unreachable; rejecting it keeps a future caller + that passes a foreign plan from failing on a tuple unpack instead. + """ + meta = self._layer_group_ranges(plan, self.pool_layers[plan.name]) + if meta is None: + raise ValueError( + f"UMBP pool {plan.name} covers none of its own layers " + f"({self.pool_layers[plan.name]})." + ) + return meta + + @staticmethod + def _plans_covering( + by_layer: dict[int, list[_PoolRangePlan]], group: list[int] + ) -> list[_PoolRangePlan]: + """Plans touching any layer of the group, each listed once, in order.""" + seen: set[int] = set() + plans = [] + for logical_layer in group: + for plan in by_layer.get(logical_layer, ()): + if id(plan) in seen: + continue + seen.add(id(plan)) + plans.append(plan) + return plans + + def _range_items(self, plan: _PoolRangePlan, layers: list[int]): + """(base_ptr, row_stride, size, offset) per emitted range, in wire order. + + Grouped by layer, and within a layer by component. Every object emits + this same tuple sequence; the only thing that varies across objects is + the pointer, by ``row * row_stride``. That invariant is what the + vectorized builder rests on. + """ + entry = self.pools[plan.name] + items: list[list[tuple[int, int, int, int]]] = [] + for logical_layer in layers: + buffer_index = entry.layer_mapping.get(logical_layer) + if buffer_index is None: + continue + items.append( + [ + (*component[buffer_index], offsets[buffer_index]) + for component, offsets in zip( + entry.buffer_meta, entry._component_offsets + ) + ] + ) + return items + + def _layer_group_ranges(self, plan: _PoolRangePlan, layers: list[int]): + """One group of layers' ranges, one nested list per object. + + Built column-wise rather than object by object. A 256K restore emits + ~63k ranges across its pools, and assembling those lists one page at a + time cost ~43 ms per load -- serialized ahead of every group's transfer, + on the same thread, so it landed directly on TTFT. Two redundancies pay + for that: ``sizes`` and ``offsets`` do not depend on the row yet were + rebuilt for every page, and ``ptrs`` is affine in the row so the whole + column can be computed at once. + + Moving the work to a helper thread was tried and does not pay: the load + thread has to re-acquire the GIL between blocking transfers and gives the + saving straight back. It has to go away rather than move. + + Returns None when this pool covers none of the group, so the caller can + skip the call. + """ + items = self._range_items(plan, layers) + if not items: + return None + rows = np.asarray(plan.locations, dtype=np.int64) + if not rows.size: + return [], [], [] + expected = len(rows) * plan.entries_per_page + if expected != len(plan.keys): + # One object per key, so a mismatch means pointers would be paired + # with the wrong objects rather than anything failing loudly. + raise ValueError( + f"UMBP pool {plan.name} has {len(plan.keys)} keys for " + f"{len(rows)} rows at {plan.entries_per_page} per page." + ) + + if self.pools[plan.name].packed: + # One object per page, its ranges running (layer, component). + flat = [item for layer_items in items for item in layer_items] + base = np.fromiter((item[0] for item in flat), np.int64, len(flat)) + stride = np.fromiter((item[1] for item in flat), np.int64, len(flat)) + ptrs = (rows[:, None] * stride[None, :] + base[None, :]).tolist() + # One list shared by every object instead of a copy each: the client + # only reads these. + sizes = [item[2] for item in flat] + offsets = [item[3] for item in flat] + return ptrs, [sizes] * len(rows), [offsets] * len(rows) + + # One object per (page, component), its ranges running over the layers. + components = len(items[0]) + base = np.array([[i[0] for i in layer] for layer in items], np.int64).T + stride = np.array([[i[1] for i in layer] for layer in items], np.int64).T + ptrs = ( + (rows[:, None, None] * stride[None, :, :] + base[None, :, :]) + .reshape(len(rows) * components, -1) + .tolist() + ) + sizes = [[layer[index][2] for layer in items] for index in range(components)] + offsets = [[layer[index][3] for layer in items] for index in range(components)] + return ptrs, sizes * len(rows), offsets * len(rows) + + @staticmethod + def _entries_per_call(sizes: list[list[int]]) -> int: + """Objects per RPC, budgeted by the ranges they actually carry. + + Counted from the ranges that were built, not from the layer count. A + packed pool puts one range per component per layer on an object, so a + packed K/V pool carries twice what the layer count suggests and the + budget would be overshot by that factor. + """ + ranges_per_object = max((len(entry) for entry in sizes), default=1) + return max(1, RANGES_PER_CALL // max(1, ranges_per_object)) + + def _run_layer_wise_batch( + self, counter_index: int, plans: list[_PoolRangePlan], ready_event: object + ) -> None: + try: + ready_event.synchronize() + by_layer: dict[int, list[_PoolRangePlan]] = defaultdict(list) + for plan in plans: + for logical_layer in self.pool_layers[plan.name]: + by_layer[logical_layer].append(plan) + + for group in self._layer_groups(): + for plan in self._plans_covering(by_layer, group): + meta = self._layer_group_ranges(plan, group) + if meta is None: + continue + ptrs, sizes, offsets = meta + step = self._entries_per_call(sizes) + for start in range(0, len(plan.keys), step): + end = start + step + chunk_keys = plan.keys[start:end] + results = list( + self.storage.client.batch_get_ranges_into_ptr( + chunk_keys, + ptrs[start:end], + sizes[start:end], + offsets[start:end], + ) + ) + if len(results) != len(chunk_keys) or not all(results): + where = ( + f"layer={group[0]}" + if len(group) == 1 + else f"layers={group[0]}..{group[-1]}" + ) + raise RuntimeError( + f"UMBP get failed for pool={plan.name}, {where}: " + f"success={sum(bool(value) for value in results)}/" + f"{len(chunk_keys)}." + ) + # Only now is every layer in the group readable, so they are + # released together. A group wider than 1 trades overlap + # granularity for fewer times each object is named on the wire. + for logical_layer in group: + self.layer_done_counter.complete(counter_index, logical_layer) + except BaseException as error: + self.layer_done_counter.fail(counter_index, error) + logger.exception("UMBP layer-wise load batch failed") + + def _layer_groups(self) -> list[list[int]]: + return [ + list(range(start, min(start + self.layer_group, self.num_layers))) + for start in range(0, self.num_layers, self.layer_group) + ] + + def offload(self, transfers: list[PoolTransfer]) -> bool: + expanded = self.pool_group.resolve_transfers(transfers, allow_partial=True) + if not expanded: + return False + self._freeze_gc_once() + ready_event = device_module.Event() + ready_event.record() + self._offload_queue.put((expanded, ready_event)) + return True + + def _take_offload_batch(self) -> tuple[list[_OffloadTask], bool]: + """Block for one task, then take whatever else is already queued. + + Returns the tasks and whether the stop sentinel came with them; each + item taken here needs one ``task_done()`` from the caller. Taking only + what is already queued keeps a task from waiting on an unsubmitted peer. + """ + first = self._offload_queue.get() + if first is None: + return [], True + tasks = [first] + pages = _offload_task_pages(first[0]) + while pages < self._offload_coalesce_pages: + try: + task = self._offload_queue.get_nowait() + except Empty: + break + if task is None: + return tasks, True + tasks.append(task) + pages += _offload_task_pages(task[0]) + return tasks, False + + def _offload_thread_func(self) -> None: + while True: + tasks, stopping = self._take_offload_batch() + taken = len(tasks) + int(stopping) + try: + if tasks: + self._offload_batch(tasks) + finally: + for _ in range(taken): + self._offload_queue.task_done() + if stopping: + return + + def _offload_batch(self, tasks: list[_OffloadTask]) -> None: + success = False + try: + success = self._run_offload(tasks) + except BaseException: + logger.exception("UMBP offload failed") + success = False + finally: + # One result per task in submission order: the tree pairs them + # positionally. A batch resolves as a unit because the first failed + # pool stops the rest, leaving every task in it incomplete. + for _ in tasks: + self._offload_results.put(success) + + def _run_offload(self, tasks: list[_OffloadTask]) -> bool: + for _, ready_event in tasks: + ready_event.synchronize() + next_slot = 0 + + def materialize_indices(indices: torch.Tensor) -> torch.Tensor: + nonlocal next_slot + slot = next_slot + next_slot += 1 + return self._materialize_offload_indices(indices, slot) + + plans = self._build_load_plans( + [expanded for expanded, _ in tasks], + materialize_indices=materialize_indices, + ) + # Over plans, not transfers: a plan already carries every task's keys + # for its pool, so walking transfers would put that pool once per task. + for plan in plans: + entry = self.pools[plan.name] + # From the pool layout, never from the ranges below: see + # _object_sizes_per_page. + per_page = _object_sizes_per_page(entry) + if len(per_page) != plan.entries_per_page: + raise ValueError( + f"UMBP pool {plan.name} declares {len(per_page)} object " + f"sizes per page but yields {plan.entries_per_page} objects." + ) + object_sizes = [ + per_page[index % plan.entries_per_page] + for index in range(len(plan.keys)) + ] + ptrs, sizes, offsets = self._all_layer_ranges(plan) + + # Preserve page order so the tier can collapse a layer into a strided copy. + + # An object's ranges must tile it exactly, so a chunk boundary may + # fall between objects but never inside one. + step = self._entries_per_call(sizes) + for start in range(0, len(plan.keys), step): + end = start + step + chunk_keys = plan.keys[start:end] + results = list( + self.storage.client.batch_put_ranges_from_ptr( + chunk_keys, + object_sizes[start:end], + ptrs[start:end], + sizes[start:end], + offsets[start:end], + ) + ) + if len(results) != len(chunk_keys) or not all(results): + logger.warning( + "UMBP offload failed: pool=%s object_range=[%d,%d) " + "success=%d/%d returned=%d", + plan.name, + start, + min(end, len(plan.keys)), + sum(bool(value) for value in results), + len(chunk_keys), + len(results), + ) + return False + + self._stats["offload"] += len(tasks) + self._stats["offload_batches"] += 1 + return True + + def _freeze_gc_once(self) -> None: + if self._gc_frozen: + return + freeze_gc("UMBP direct linker") + self._gc_frozen = True + + def num_completed_offloads(self) -> int: + # The tree agrees on the drain count across ranks before calling pop. + return self._offload_results.qsize() + + def pop_completed_offload(self) -> bool: + return self._offload_results.get_nowait() + + def reset(self) -> None: + self._pending.clear() + self._load_queue.join() + self._offload_queue.join() + while True: + try: + self._offload_results.get_nowait() + except Empty: + break + while True: + try: + self._completed_loads.get_nowait() + except Empty: + break + self.layer_done_counter.reset() + + def close(self) -> None: + if self._closed: + return + self.reset() + for thread, queue in ( + (self._offload_thread, self._offload_queue), + (self._load_thread, self._load_queue), + ): + if thread.is_alive(): + queue.put(None) + thread.join() + if self._standalone_process_mode and self._registered: + # StandaloneProcess deregistration is client-wide; one call tears + # down every registered region. Keep the GPU tensors alive until + # the synchronous RPC has completed successfully. + self.storage.client.deregister_memory(self._registered[0][0]) + logger.info("UMBP direct linker stats: %s", self._stats) + self.storage.close() + self._closed = True 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 index fef8ecb3d..d903d294e 100644 --- a/python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_host_allocator.py @@ -45,6 +45,9 @@ class UMBPHostTensorAllocator(HostTensorAllocator): ) self._numa_node = _int_env("SGLANG_HICACHE_HOST_NUMA_NODE", -1) self._prefault = _bool_env("SGLANG_HICACHE_HOST_PREFAULT", True) + # Standalone mode needs fd-shareable backing; allocation precedes + # config parsing. + self._standalone_process = bool(os.getenv("UMBP_STANDALONE_ADDRESS")) self._handles: Dict[int, Any] = {} def allocate( @@ -62,11 +65,18 @@ class UMBPHostTensorAllocator(HostTensorAllocator): 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 - ) + if self._standalone_process: + requested_backing = ( + self._mod.UMBPHostBufferBacking.AnonymousShmHugetlb + if self._use_hugepage + else self._mod.UMBPHostBufferBacking.AnonymousShm + ) + else: + requested_backing = ( + self._mod.UMBPHostBufferBacking.AnonymousHugetlb + if self._use_hugepage + else self._mod.UMBPHostBufferBacking.Anonymous + ) handle = self._allocator.alloc( nbytes, @@ -101,15 +111,19 @@ class UMBPHostTensorAllocator(HostTensorAllocator): handle.mapped_size, self._numa_node, ) - if ( - self._use_hugepage - and handle.actual_backing == self._mod.UMBPHostBufferBacking.Anonymous - ): + demoted = handle.actual_backing == ( + self._mod.UMBPHostBufferBacking.AnonymousShm + if self._standalone_process + else self._mod.UMBPHostBufferBacking.Anonymous + ) + if self._use_hugepage and demoted: 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." + "UMBPHostTensorAllocator: requested %s backing but kernel " + "demoted to %s (4 KiB pages). Check vm.nr_hugepages and " + "HugePages_Free in /proc/meminfo. Performance and AINIC " + "MR-size benefits will not apply.", + requested_backing, + handle.actual_backing, ) return tensor.view(dims) diff --git a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py index c6affd41b..eed41159d 100644 --- a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py @@ -38,6 +38,8 @@ def _import_umbp_client(): UMBPIoBackend = getattr(umbp_mod, "UMBPIoBackend", None) UMBPDurabilityMode = getattr(umbp_mod, "UMBPDurabilityMode", None) UMBPDistributedConfig = getattr(umbp_mod, "UMBPDistributedConfig", None) + UMBPStandaloneProcessConfig = getattr(umbp_mod, "UMBPStandaloneProcessConfig", None) + UMBPDeploymentMode = getattr(umbp_mod, "UMBPDeploymentMode", None) return ( UMBPClient, @@ -46,6 +48,8 @@ def _import_umbp_client(): UMBPIoBackend, UMBPDurabilityMode, UMBPDistributedConfig, + UMBPStandaloneProcessConfig, + UMBPDeploymentMode, ) @@ -134,6 +138,10 @@ def _select_rank_config_value( # knobs outside this list go through the "spdk_passthrough" escape hatch. _COMMON_EXTRA_KEYS = frozenset( { + "node_address", + "node_id", + "node_tags", + "tags", "dram_capacity_bytes", "ssd_enabled", "ssd_storage_dir", @@ -164,6 +172,8 @@ _COMMON_EXTRA_KEYS = frozenset( "kv_events_subscriber", "kv_events_endpoint", "kv_events_topic", + "disable_zero_copy_register", + "extra_backend_tag", } ) @@ -175,24 +185,25 @@ _STANDALONE_ONLY_EXTRA_KEYS = frozenset( "eviction_policy", "eviction_candidate_window", "auto_promote_on_read", + "standalone_address", + "standalone_auto_start", + "standalone_startup_timeout_ms", } ) _DISTRIBUTED_ONLY_EXTRA_KEYS = frozenset( { "master_address", - "node_address", - "node_id", "auto_heartbeat", "io_engine_host", "io_engine_port", "staging_buffer_size", + "ranged_scratch_size", "ssd_staging_buffer_size", "ssd_staging_buffer_slots", "peer_service_port", "cache_remote_fetches", "dram_page_size", - "disable_zero_copy_register", } ) @@ -238,7 +249,13 @@ class UMBPStore(HiCacheStorage): self, storage_config: HiCacheStorageConfig = None, mem_pool_host: HostKVCache = None, + *, + per_rank_keyspace: bool = False, ): + # per_rank_keyspace: the direct linker already organises keys by its own + # cache-group rank, so HiCache's shared-SSD leader/follower deduplication + # must not be layered on top. Default False preserves the existing + # HiCache L3 behaviour. ( UMBPClient, UMBPConfig, @@ -246,6 +263,8 @@ class UMBPStore(HiCacheStorage): UMBPIoBackend, UMBPDurabilityMode, UMBPDistributedConfig, + UMBPStandaloneProcessConfig, + UMBPDeploymentMode, ) = _import_umbp_client() if storage_config is not None: @@ -260,6 +279,7 @@ class UMBPStore(HiCacheStorage): self.pp_rank = 0 self.pp_size = 1 self.tp_size = 1 + self._umbp_deployment_mode_enum = UMBPDeploymentMode cfg = UMBPConfig.from_environment() # UMBPStore owns role selection explicitly. Do not inherit LOCAL_RANK / @@ -268,6 +288,12 @@ class UMBPStore(HiCacheStorage): # and skip writes. cfg.role = UMBPRole.Standalone extra = getattr(storage_config, "extra_config", None) or {} + prefix_parts = [] + if extra.get("extra_backend_tag") is not None: + prefix_parts.append(str(extra["extra_backend_tag"])) + if storage_config is not None and storage_config.model_name: + prefix_parts.append("-".join(storage_config.model_name.split("/"))) + self.config_prefix = "_".join(prefix_parts) if prefix_parts else None explicit_tenant_id = ( os.getenv("UMBP_SPDK_PROXY_TENANT_ID") is not None or "spdk_proxy_tenant_id" in extra @@ -461,6 +487,58 @@ class UMBPStore(HiCacheStorage): master_address = extra.get( "master_address", _optional_env_str("UMBP_MASTER_ADDRESS") ) + standalone_extra_address = extra.get("standalone_address") + standalone_env_address = _optional_env_str("UMBP_STANDALONE_ADDRESS") + standalone_address = standalone_extra_address or standalone_env_address + # Verify the client did not silently fall back to local mode. + self._standalone_process_expected = bool(standalone_address) + if master_address and standalone_address: + raise ValueError( + "master_address and standalone_address are mutually exclusive " + "(distributed vs. standalone-process mode)." + ) + if ( + mem_pool_host is not None + and standalone_extra_address + and not standalone_env_address + ): + raise ValueError( + "standalone_address in hicache-storage-backend-extra-config is " + "not supported when a host KV pool is present. The host memory " + "pool allocator chooses Anonymous vs. AnonymousShm before " + "extra_config is parsed, so set UMBP_STANDALONE_ADDRESS in the " + "process environment instead." + ) + + # Both remote modes use the same worker identity. + def _resolve_node_address() -> str: + node_address = extra.get( + "node_address", _optional_env_str("UMBP_NODE_ADDRESS") + ) + if node_address is None: + return _default_node_address() + return _select_rank_config_value( + node_address, unique_rank, "node_address", str + ) + + def _resolve_node_id(node_address: str) -> str: + node_id = extra.get("node_id", _optional_env_str("UMBP_NODE_ID")) + if node_id is None: + return ( + 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}" + ) + return _select_rank_config_value(node_id, unique_rank, "node_id", str) + + def _resolve_node_tags() -> List[str]: + raw_tags = extra.get("node_tags", extra.get("tags")) + if raw_tags is None: + raw_tags = _optional_env_str("UMBP_NODE_TAGS") + if raw_tags is None: + return [] + if isinstance(raw_tags, str): + return [tag.strip() for tag in raw_tags.split(",") if tag.strip()] + return [str(tag) for tag in raw_tags] _warn_extra_config_scope(extra, distributed_enabled=bool(master_address)) if master_address and UMBPDistributedConfig is not None: @@ -470,33 +548,11 @@ class UMBPStore(HiCacheStorage): 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, - ) + node_address = _resolve_node_address() 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, - ) + dist_cfg.master_config.node_id = _resolve_node_id(node_address) + if hasattr(dist_cfg.master_config, "tags"): + dist_cfg.master_config.tags = _resolve_node_tags() if "auto_heartbeat" in extra: dist_cfg.master_config.auto_heartbeat = _strict_bool( @@ -531,6 +587,10 @@ class UMBPStore(HiCacheStorage): if "staging_buffer_size" in extra: dist_cfg.staging_buffer_size = int(extra["staging_buffer_size"]) + if "ranged_scratch_size" in extra and hasattr( + dist_cfg, "ranged_scratch_size" + ): + dist_cfg.ranged_scratch_size = int(extra["ranged_scratch_size"]) if "ssd_staging_buffer_size" in extra and hasattr( dist_cfg, "ssd_staging_buffer_size" @@ -604,8 +664,8 @@ class UMBPStore(HiCacheStorage): meta = mem_pool_host.get_split_heads_page_buffer_meta(dummy, sf) else: meta = mem_pool_host.get_page_buffer_meta(dummy) - # meta is None for a logical-anchor group (see note above); - # esz is the per-page element-size list otherwise. + # A hybrid logical anchor returns None here by design; leave + # dram_page_size at 0 and let the per-pool v2 sizes handle it. esz = meta[1] if meta else None page_byte_size = int(esz[0]) if esz else 0 @@ -647,6 +707,52 @@ class UMBPStore(HiCacheStorage): dist_cfg.io_engine.port, dist_cfg.peer_service_port, ) + elif standalone_address: + if UMBPStandaloneProcessConfig is None: + raise RuntimeError( + "Installed mori does not expose UMBPStandaloneProcessConfig" + ) + standalone_cfg = UMBPStandaloneProcessConfig() + standalone_cfg.address = str(standalone_address) + auto_start = extra.get( + "standalone_auto_start", + _optional_env_str("UMBP_STANDALONE_AUTO_START"), + ) + if auto_start is not None: + standalone_cfg.auto_start = _strict_bool( + auto_start, "standalone_auto_start" + ) + startup_timeout_ms = extra.get( + "standalone_startup_timeout_ms", + _optional_env_int("UMBP_STANDALONE_STARTUP_TIMEOUT_MS"), + ) + if startup_timeout_ms is not None: + standalone_cfg.startup_timeout_ms = int(startup_timeout_ms) + if standalone_cfg.startup_timeout_ms <= 0: + raise ValueError("standalone_startup_timeout_ms must be > 0") + if all( + hasattr(standalone_cfg, field) + for field in ("worker_node_address", "worker_node_id", "tags") + ): + worker_node_address = _resolve_node_address() + standalone_cfg.worker_node_address = worker_node_address + standalone_cfg.worker_node_id = _resolve_node_id(worker_node_address) + standalone_cfg.tags = _resolve_node_tags() + else: + logger.warning( + "UMBPStore standalone-process mode: installed mori does not " + "expose worker identity on UMBPStandaloneProcessConfig; a " + "distributed-backed standalone server cannot build per-worker " + "external-KV identities." + ) + cfg.standalone_process = standalone_cfg + logger.info( + "UMBPStore standalone-process mode: address=%s, auto_start=%s, " + "startup_timeout_ms=%s", + standalone_cfg.address, + standalone_cfg.auto_start, + standalone_cfg.startup_timeout_ms, + ) self.storage_config = storage_config @@ -657,8 +763,25 @@ class UMBPStore(HiCacheStorage): 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: + remote_process_enabled = ( + cfg.distributed is not None or cfg.standalone_process is not None + ) + # Shared SSD exists to deduplicate MLA KV, which TP replicates: one + # rank owns the bytes and the others read them back. A caller that + # already keys per rank has nothing to deduplicate -- and would be + # broken by the scheme, because a follower would be sent looking for + # keys the leader never wrote under the follower's own suffix. + # + # It is also the reason embedded mode could not run: followers are the + # one role whose client reports no ranged multi-buffer I/O, which + # page-granular objects require. Standalone never hit this, not by + # design but because remote_process_enabled short-circuits it there. + if ( + not remote_process_enabled + and not per_rank_keyspace + 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. @@ -802,11 +925,10 @@ class UMBPStore(HiCacheStorage): safe_cap = int(cfg.ssd.capacity_bytes * 0.95) cfg.ssd.spdk_proxy_tenant_quota_bytes = max(1, safe_cap // dp_size_hint) - # Initialize registration state before the optional constructor-time - # register_mem_pool_host() call below. In particular, do not overwrite - # the logical-anchor flag after that call has detected a LogicalHostPool. + # Initialize before the optional constructor-time pool registration. self.registered_pools: dict = {} self._kv_anchor_is_logical = False + self._registered_regions: set = set() self.client = UMBPClient(cfg) if mem_pool_host is not None: @@ -888,45 +1010,99 @@ class UMBPStore(HiCacheStorage): "page_head", ], "UMBP store only supports page_first, page_first_direct, or page_head layout" - # Hybrid logical anchors (e.g. DeepSeek-V4's KV anchor LogicalHostPool) - # own only allocation indices and hold no physical KV tensor. Compute - # this once and reuse: there is nothing to register for RDMA here, v1 - # I/O no-ops on it, and the real per-pool buffers are registered through - # register_mem_host_pool_v2(). + # A logical anchor owns indices; side pools carry the data. self._kv_anchor_is_logical = self.mem_pool_host.kv_buffer is None - self._zero_copy_registered = False + + # Side-pool registration needs the mode even for a logical anchor. + self._is_standalone_process = False + if self.client is not None: + deployment_mode = None + mode_enum = self._umbp_deployment_mode_enum + try: + deployment_mode = self.client.get_deployment_mode() + if mode_enum is not None: + self._is_standalone_process = ( + deployment_mode == mode_enum.StandaloneProcess + ) + except Exception as exc: + if self._standalone_process_expected: + raise RuntimeError( + "UMBPStore expected standalone-process mode from " + "UMBP_STANDALONE_ADDRESS, but get_deployment_mode() failed." + ) from exc + if self._standalone_process_expected: + if mode_enum is None: + raise RuntimeError( + "UMBPStore expected standalone-process mode, but " + "UMBPDeploymentMode is not exposed by mori.umbp." + ) + if deployment_mode != mode_enum.StandaloneProcess: + raise RuntimeError( + "UMBPStore expected standalone-process mode, but the " + f"UMBP client reported deployment_mode={deployment_mode!r}." + ) + if self._kv_anchor_is_logical: return - # 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. - if self._register_host_buffer_for_zero_copy(mem_pool_host): - self._zero_copy_registered = True + self._zero_copy_registered = self._register_host_buffer_for_zero_copy( + mem_pool_host + ) + + @staticmethod + def _pool_physical_buffers(host_pool: HostKVCache) -> List[Any]: + """Return every non-empty physical tensor exposed by a host pool.""" + getter = getattr(host_pool, "get_hybrid_pool_buffer", None) + buffers = getter() if getter is not None else None + if not buffers: + buffers = [getattr(host_pool, "kv_buffer", None)] + flat: List[Any] = [] + for buffer in buffers: + if buffer is None: + continue + for tensor in buffer if isinstance(buffer, (list, tuple)) else [buffer]: + # Empty views may share storage; register_memory rejects zero bytes. + if tensor is not None and tensor.numel() > 0: + flat.append(tensor) + return flat + + @staticmethod + def _buffer_extent(buffer, allocator) -> tuple: + """Registerable (base pointer, size) of the allocation behind a tensor.""" + storage = buffer.untyped_storage() + base = int(storage.data_ptr()) + size = int(storage.nbytes()) + # Hugepage-backed mmaps are rounded up to the hugepage boundary, and + # ibv_reg_mr on AINIC / ROCm needs whole hugepages covered. + mapped_size_fn = getattr(allocator, "mapped_size_for", None) + mapped_size = ( + mapped_size_fn(base) + if mapped_size_fn is not None + else getattr(allocator, "mapped_size", 0) + ) + return base, max(size, int(mapped_size or 0)) def _register_host_buffer_for_zero_copy(self, host_pool: HostKVCache) -> bool: - """Register a host pool's KV buffer with the RDMA IOEngine for zero-copy. - - Shared by the single-pool path (register_mem_pool_host) and the - multi-pool path (register_mem_host_pool_v2). Returns True when the - buffer was successfully registered, False on any skip/failure (the - caller then transparently falls back to the staging-buffer path). - """ + """Register host buffers; standalone failures are fatal without fallback.""" if self.client is None: return False + is_standalone_process = getattr(self, "_is_standalone_process", False) try: is_distributed = bool(self.client.is_distributed()) except Exception: is_distributed = False - if not is_distributed: + if not (is_distributed or is_standalone_process): return False if not hasattr(self.client, "register_memory"): return False if getattr(self, "_disable_zero_copy_register", False): + if is_standalone_process: + raise RuntimeError( + "disable_zero_copy_register is not supported in UMBP " + "standalone-process mode: there is no staging-buffer " + "fallback path." + ) logger.info( "UMBPStore: skipping host KV buffer RDMA registration because " "disable_zero_copy_register=true (UMBP_DISABLE_ZERO_COPY_REGISTER). " @@ -934,69 +1110,62 @@ class UMBPStore(HiCacheStorage): "size is capped by distributed.staging_buffer_size." ) return False - # NOTE(layer_first): this only handles the page_first layout, where a - # host pool exposes a single contiguous `kv_buffer` that we can register - # for RDMA in one shot. If UMBP later supports a layer_first layout, or - # side pools that expose multiple buffers via get_hybrid_pool_buffer() - # (e.g. DSAIndexerPoolHost, whose buffer lives in - # index_k_with_scale_buffer rather than kv_buffer), this branch must be - # extended to register every per-layer / per-buffer region. Otherwise - # such pools bypass zero-copy and silently fall back to the slower - # staging-buffer path. - kv_buffer = getattr(host_pool, "kv_buffer", None) - if kv_buffer is None: + buffers = self._pool_physical_buffers(host_pool) + if not buffers: + if is_standalone_process: + raise RuntimeError( + f"UMBPStore: {type(host_pool).__name__} exposes no host buffer " + "to register; standalone-process mode has no fallback path." + ) return False - try: - 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(host_pool, "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 False - if ok: + + mode = "standalone-process" if is_standalone_process else "distributed" + allocator = getattr(host_pool, "allocator", None) + # Already registered storage counts as covered. + covered = 0 + for buffer in buffers: + try: + host_ptr, host_size = self._buffer_extent(buffer, allocator) + if host_ptr in self._registered_regions: + covered += 1 + continue + ok = bool(self.client.register_memory(host_ptr, host_size)) + except Exception as exc: + if is_standalone_process: + raise RuntimeError( + "UMBPStore: register_memory failed in standalone-process " + f"mode and cannot fall back: {exc}" + ) from 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 False + if not ok: + if is_standalone_process: + raise RuntimeError( + "UMBPStore: register_memory returned false in " + "standalone-process mode; no fallback path exists." + ) + logger.warning( + "UMBPStore: register_memory returned false; staying on staging " + "buffer fallback path." + ) + return False + self._registered_regions.add(host_ptr) + covered += 1 logger.info( - "UMBPStore: registered host KV buffer for RDMA zero-copy " - "(ptr=0x%x, size=%d MB)", + "UMBPStore: registered host buffer for zero-copy " + "(ptr=0x%x, size=%d MB, mode=%s)", host_ptr, host_size // (1024 * 1024), + mode, ) - return True - logger.warning( - "UMBPStore: register_memory returned false; staying on staging " - "buffer fallback path." - ) - return False + return covered == len(buffers) def register_mem_host_pool_v2(self, host_pool: HostKVCache, host_pool_name): - """Register an additional hybrid side pool (DeepSeek-V4 HostPoolGroup). - - The controller calls this once per PoolEntry in the group, including the - KV anchor. The KV anchor is logical (no physical tensor) so we skip it; - its allocation-index role is unrelated to storage I/O. Every other pool - (SWA / compressed KV / indexer / state) carries a real page_first KV - buffer that must be (a) resolvable by name at v2 I/O time and (b) - registered with the RDMA IOEngine for zero-copy transfers. - """ - # KV anchor is either already registered via register_mem_pool_host() - # (non-hybrid single pool) or purely logical (hybrid group). Skip it. if host_pool_name == PoolName.KV: return self.registered_pools[host_pool_name] = host_pool @@ -1078,8 +1247,6 @@ class UMBPStore(HiCacheStorage): extra_info: Optional[HiCacheStorageExtraInfo] = None, ) -> List[bool]: if self._kv_anchor_is_logical: - # DeepSeek-V4's KV anchor is logical only; the physical KV data is - # carried by the v2 side pools, so there is nothing to read here. return [True] * len(keys) key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices) @@ -1152,8 +1319,6 @@ class UMBPStore(HiCacheStorage): return [True] * page_count if self._kv_anchor_is_logical: - # DeepSeek-V4's KV anchor is logical only; the physical KV data is - # written by the v2 side pools, so there is nothing to write here. return [True] * len(keys) key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices) @@ -1219,47 +1384,66 @@ class UMBPStore(HiCacheStorage): return hit_count // key_multiplier # ------------------------------------------------------------------ - # Multi-pool v2 interface (DeepSeek-V4 hybrid HiCache HostPoolGroup) - # - # The DeepSeek-V4 HiCache stack splits KV state across several page_first - # side pools (SWA / compressed KV / indexer / state), coordinated by a - # logical KV anchor that owns only page indices. The controller registers - # each real pool through register_mem_host_pool_v2() and drives storage - # via these _v2 methods, one PoolTransfer per pool. This mirrors the proven - # MooncakeStore / HiCacheHF3FS design, specialized for UMBP's page_first, - # single-object-per-page layout (each page -> exactly one storage object). + # Multi-pool v2 interface # ------------------------------------------------------------------ - def _get_hybrid_page_component_keys(self, page_keys, transfer: PoolTransfer): - """Map per-page logical keys to per-object storage keys for a side pool. - - For UMBP every registered side pool is page_first and stores one object - per page (MLA: a single K object; MHA: a K and a V object), so the - component-key count is an exact multiple of the page count. The pool - name is embedded in the suffix so pages that share a hash across pools - never collide. - """ + def _get_hybrid_page_component_keys( + self, page_keys, transfer: PoolTransfer, *, rank_suffix: Optional[str] = None + ): + """Expand logical page keys for one registered hybrid side pool.""" pool_name = transfer.name host_pool = self.registered_pools.get(pool_name) if host_pool is None: raise ValueError(f"Unregistered UMBP hybrid pool: {pool_name}") - if self.is_mla_backend: - # Single compressed object per page. - suffixes = [f"_{self.mla_suffix}_{pool_name}"] + mla_suffix = self.mla_suffix if rank_suffix is None else rank_suffix + mha_suffix = ( + getattr(self, "mha_suffix", mla_suffix) + if rank_suffix is None + else rank_suffix + ) + + components = getattr(host_pool, "components", None) + if pool_name == PoolName.MAMBA: + conv_num = len(getattr(host_pool, "conv_buffer", None) or []) + suffixes = [f"_{mha_suffix}_conv_{i}" for i in range(conv_num)] + if getattr(host_pool, "temporal_state_elem_size", 1) > 0: + suffixes = [f"_{mha_suffix}_temporal"] + suffixes + elif components is not None and len(components) == 1: + suffixes = [f"_{mla_suffix}_{pool_name}"] + elif components is not None and len(components) == 2: + # Packed DevicePoolEntry K/V components share one stored object. + suffixes = ( + [f"_{mha_suffix}_{pool_name}"] + if host_pool.packed + else [ + f"_{mha_suffix}_{pool_name}_k", + f"_{mha_suffix}_{pool_name}_v", + ] + ) + elif components is not None: + raise ValueError( + f"Unsupported UMBP component count for pool {pool_name}: " + f"{len(components)}" + ) + elif self.is_mla_backend: + suffixes = [f"_{mla_suffix}_{pool_name}"] elif getattr(host_pool, "v_buffer", None) is not None: - # Ordinary MHA side pool mirrors a K/V pool. suffixes = [ - f"_{self.mha_suffix}_{pool_name}_k", - f"_{self.mha_suffix}_{pool_name}_v", + f"_{mha_suffix}_{pool_name}_k", + f"_{mha_suffix}_{pool_name}_v", ] else: - suffixes = [f"_{self.mha_suffix}_{pool_name}"] + suffixes = [f"_{mha_suffix}_{pool_name}"] - key_multiplier = len(suffixes) component_keys = [ f"{page_key}{suffix}" for page_key in page_keys for suffix in suffixes ] - return component_keys, key_multiplier + if self.config_prefix: + component_keys = [ + f"{self.config_prefix}_{component_key}" + for component_key in component_keys + ] + return component_keys, len(suffixes) def batch_exists_v2( self, @@ -1267,20 +1451,18 @@ class UMBPStore(HiCacheStorage): pool_transfers: Optional[List[PoolTransfer]] = None, extra_info: Optional[HiCacheStorageExtraInfo] = None, ) -> PoolTransferResult: - if self._kv_anchor_is_logical: - # Logical KV anchor: no physical KV object exists in UMBP, so the - # usable prefix is bounded entirely by the required side pools. - kv_pages = len(keys) - else: - kv_pages = self.batch_exists(keys, extra_info) - + kv_pages = ( + len(keys) + if self._kv_anchor_is_logical + else self.batch_exists(keys, extra_info) + ) hit_count: dict = {PoolName.KV: kv_pages} if kv_pages else {} final_pages = kv_pages for transfer in pool_transfers or []: if final_pages == 0: break - component_keys, key_multiplier = self._get_hybrid_page_component_keys( + component_keys, multiplier = self._get_hybrid_page_component_keys( keys[:final_pages], transfer ) exists = list(self.client.batch_exists(component_keys)) @@ -1294,18 +1476,15 @@ class UMBPStore(HiCacheStorage): ) final_pages = 0 break - # Collapse per-object results into per-page presence. page_exists = [ - all(exists[i * key_multiplier : (i + 1) * key_multiplier]) + all(exists[i * multiplier : (i + 1) * multiplier]) for i in range(final_pages) ] - boundary = 0 if transfer.hit_policy == PoolHitPolicy.ALL_PAGES: - try: - boundary = page_exists.index(False) - except ValueError: - boundary = final_pages + boundary = ( + page_exists.index(False) if False in page_exists else final_pages + ) elif transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES: trailing = max(1, len(transfer.keys) if transfer.keys else 1) for prefix_len in range(final_pages, 0, -1): @@ -1334,38 +1513,26 @@ class UMBPStore(HiCacheStorage): if not keys or host_indices is None: results[transfer.name] = [False] * len(keys) continue - assert len(keys) == len(host_indices) // page_size + if len(keys) != len(host_indices) // page_size: + raise ValueError( + f"UMBP v2 pool {transfer.name} has {len(keys)} keys for " + f"{len(host_indices)} indices with page_size={page_size}." + ) - key_strs, key_multiplier = self._get_hybrid_page_component_keys( - keys, transfer + key_strs, multiplier = self._get_hybrid_page_component_keys(keys, transfer) + ptrs, sizes = host_pool.get_page_buffer_meta(host_indices) + if not len(key_strs) == len(ptrs) == len(sizes): + raise ValueError( + f"UMBP v2 buffer-meta mismatch for pool {transfer.name}: " + f"keys={len(key_strs)} ptrs={len(ptrs)} sizes={len(sizes)}" + ) + + operation = ( + self.client.batch_put_from_ptr + if is_set + else self.client.batch_get_into_ptr ) - ptr_list, element_size_list = host_pool.get_page_buffer_meta(host_indices) - # page_first side pools emit exactly one (ptr, size) per component - # key; assert the invariant so any future layout change is caught - # loudly instead of silently corrupting the key<->buffer zip. - assert len(key_strs) == len(ptr_list) == len(element_size_list), ( - f"UMBP v2 buffer-meta mismatch for pool {transfer.name}: " - f"keys={len(key_strs)} ptrs={len(ptr_list)} sizes={len(element_size_list)}" - ) - - if is_set: - # UMBP performs its own key-level deduplication, so skip the - # extra batch_exists round-trip and put directly (mirrors - # batch_set_v1). - io_results = [ - bool(r) - for r in self.client.batch_put_from_ptr( - key_strs, list(ptr_list), list(element_size_list) - ) - ] - else: - io_results = [ - bool(r) - for r in self.client.batch_get_into_ptr( - key_strs, list(ptr_list), list(element_size_list) - ) - ] - + io_results = [bool(value) for value in operation(key_strs, ptrs, sizes)] if len(io_results) != len(key_strs): logger.error( "UMBP v2 %s result-size mismatch for pool %s: " @@ -1378,9 +1545,8 @@ class UMBPStore(HiCacheStorage): results[transfer.name] = [False] * len(keys) continue - # Collapse per-object results back to per-page results. results[transfer.name] = [ - all(io_results[i * key_multiplier : (i + 1) * key_multiplier]) + all(io_results[i * multiplier : (i + 1) * multiplier]) for i in range(len(keys)) ] return results diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 40524e411..64a8d2715 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2842,7 +2842,7 @@ class ServerArgs: str, Arg( help="Storage backend for --enable-unified-cache-external-linker.", - choices=["mooncake"], + choices=["mooncake", "mori"], ), NS("memory"), ] = "mooncake" diff --git a/test/registered/hicache/test_hicache_storage_umbp_backend.py b/test/registered/hicache/test_hicache_storage_umbp_backend.py index 8145e262d..2b4e9ba08 100644 --- a/test/registered/hicache/test_hicache_storage_umbp_backend.py +++ b/test/registered/hicache/test_hicache_storage_umbp_backend.py @@ -121,6 +121,7 @@ class TestHiCacheStorageUMBPBackend(CustomTestCase): # An absent master address keeps every TP rank in standalone local mode, # so this E2E does not require an RDMA-capable CI runner. env.pop("UMBP_MASTER_ADDRESS", None) + env.pop("UMBP_STANDALONE_ADDRESS", None) env.update( { "SGLANG_ENABLE_DETERMINISTIC_INFERENCE": "1", diff --git a/test/registered/unit/mem_cache/test_registry.py b/test/registered/unit/mem_cache/test_registry.py index 5f9f903fb..cd1b52fa2 100644 --- a/test/registered/unit/mem_cache/test_registry.py +++ b/test/registered/unit/mem_cache/test_registry.py @@ -336,6 +336,57 @@ class TestDefaultRadixCacheFactory(CustomTestCase): ctx.tp_worker.register_hicache_layer_transfer_counter.assert_called_once() self.assertIs(result, fake_radix.UnifiedRadixCache.return_value) + def test_unified_radix_cache_with_mori_external_linker(self): + from sglang.srt.mem_cache.storage.umbp import umbp_direct_linker + + ctx = _make_ctx(self) + object.__setattr__( + ctx.server_args, "enable_unified_cache_external_linker", True + ) + object.__setattr__( + ctx.server_args, "unified_cache_external_linker_backend", "mori" + ) + self.assertTrue(ctx.server_args.enable_unified_cache_external_linker) + self.assertEqual(ctx.server_args.unified_cache_external_linker_backend, "mori") + fake_components = MagicMock() + fake_components.ComponentType.FULL = "full" + fake_radix = MagicMock() + cache = fake_radix.UnifiedRadixCache.return_value + cache.components = ("full",) + counter = MagicMock(name="layer_done_counter") + cache.linker.layer_done_counter = counter + linker = MagicMock(name="linker") + + with ( + patch.dict( + "sys.modules", + { + "sglang.srt.mem_cache.unified_cache.components": fake_components, + "sglang.srt.mem_cache.unified_radix_cache": fake_radix, + }, + ), + patch.object( + umbp_direct_linker, + "UMBPDirectLinker", + return_value=linker, + ) as linker_cls, + ): + result = default_radix_cache_factory(ctx) + + linker_cls.assert_called_once_with( + ctx.server_args, + ctx.params, + components={"full"}, + ) + cache.init_cache_linker.assert_called_once_with(linker) + ctx.params.token_to_kv_pool_allocator.get_kvcache.return_value.register_layer_transfer_counter.assert_called_once_with( + counter + ) + ctx.tp_worker.register_hicache_layer_transfer_counter.assert_called_once_with( + counter + ) + self.assertIs(result, cache) + def test_swa_radix_cache_when_hybrid_swa(self): ctx = _make_ctx(self, is_hybrid_swa=True) # SWA hybrid models now default to the unified radix tree. diff --git a/test/registered/unit/mem_cache/test_umbp_host_allocator.py b/test/registered/unit/mem_cache/test_umbp_host_allocator.py index 4e64b5a97..070f0d70d 100644 --- a/test/registered/unit/mem_cache/test_umbp_host_allocator.py +++ b/test/registered/unit/mem_cache/test_umbp_host_allocator.py @@ -21,6 +21,8 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu") class FakeBacking(Enum): Anonymous = 0 AnonymousHugetlb = 1 + AnonymousShm = 2 + AnonymousShmHugetlb = 3 class FakeHandle: @@ -64,7 +66,10 @@ class FakeHostMemAllocator: mapped_size=size, actual_backing=backing, actual_alignment=( - hugepage_size if backing == FakeBacking.AnonymousHugetlb else 4096 + hugepage_size + if backing + in (FakeBacking.AnonymousHugetlb, FakeBacking.AnonymousShmHugetlb) + else 4096 ), ) self.alloc_calls.append( @@ -142,6 +147,37 @@ class TestUMBPHostAllocator(unittest.TestCase): tensor.fill_(3.0) self.assertEqual(float(tensor[0, 0]), 3.0) + def test_standalone_process_uses_shareable_backing(self): + self._install_fake_mori() + + from sglang.srt.mem_cache.storage.umbp.umbp_host_allocator import ( + UMBPHostTensorAllocator, + ) + + cases = ( + ("0", FakeBacking.AnonymousShm), + ("1", FakeBacking.AnonymousShmHugetlb), + ) + for use_hugepage, expected in cases: + with ( + self.subTest(use_hugepage=use_hugepage), + mock.patch.dict( + "os.environ", + { + "UMBP_STANDALONE_ADDRESS": "unix:///tmp/umbp-test.sock", + "SGLANG_HICACHE_HOST_HUGEPAGE": use_hugepage, + }, + ), + ): + allocator = UMBPHostTensorAllocator() + tensor = allocator.allocate((16,), dtype=torch.uint8, device="cpu") + + self.assertEqual( + allocator._allocator.alloc_calls[0]["backing"], expected + ) + del tensor + allocator.__del__() + def test_umbp_allocator_del_calls_free_once(self): self._install_fake_mori() diff --git a/test/registered/unit/mem_cache/test_umbp_store.py b/test/registered/unit/mem_cache/test_umbp_store.py index 699cd3213..cede83f4e 100755 --- a/test/registered/unit/mem_cache/test_umbp_store.py +++ b/test/registered/unit/mem_cache/test_umbp_store.py @@ -119,6 +119,55 @@ def make_indices(indices): class TestUMBPStore(unittest.TestCase): + def test_standalone_process_configuration(self): + from sglang.srt.mem_cache.storage.umbp import umbp_store + + imported = list(umbp_store._import_umbp_client()) + captured = [] + + def make_client(config): + captured.append(config) + client = MagicMock() + client.flush.return_value = True + return client + + imported[0] = make_client + config = MockStorageConfig( + extra_config={ + "standalone_address": "unix:///tmp/umbp-test.sock", + "standalone_auto_start": False, + "standalone_startup_timeout_ms": 1234, + "ssd_enabled": False, + "extra_backend_tag": "tenant-a", + } + ) + with patch.object( + umbp_store, "_import_umbp_client", return_value=tuple(imported) + ): + store = umbp_store.UMBPStore(config, mem_pool_host=None) + + self.assertEqual(len(captured), 1) + self.assertIsNone(captured[0].distributed) + self.assertEqual( + captured[0].standalone_process.address, "unix:///tmp/umbp-test.sock" + ) + self.assertFalse(captured[0].standalone_process.auto_start) + self.assertEqual(captured[0].standalone_process.startup_timeout_ms, 1234) + self.assertEqual(store.config_prefix, "tenant-a_test-model") + store.close() + + def test_standalone_and_distributed_addresses_are_mutually_exclusive(self): + from sglang.srt.mem_cache.storage.umbp import umbp_store + + config = MockStorageConfig( + extra_config={ + "master_address": "127.0.0.1:1234", + "standalone_address": "unix:///tmp/umbp-test.sock", + } + ) + with self.assertRaisesRegex(ValueError, "mutually exclusive"): + umbp_store.UMBPStore(config, mem_pool_host=None) + def test_basic_set_get(self): from sglang.srt.mem_cache.storage.umbp.umbp_store import UMBPStore @@ -326,6 +375,7 @@ class TestUMBPStoreDefensiveSemantics(unittest.TestCase): store.is_mla_backend = True store.mla_suffix = "" store.mha_suffix = "0" + store.config_prefix = None store.register_mem_host_pool_v2(MockHybridSidePool(), PoolName.DEEPSEEK_V4_C4) return store @@ -345,6 +395,7 @@ class TestUMBPStoreDefensiveSemantics(unittest.TestCase): spdk_proxy_tenant_quota_bytes=0, ) self.distributed = None + self.standalone_process = None @classmethod def from_environment(cls): @@ -366,6 +417,8 @@ class TestUMBPStoreDefensiveSemantics(unittest.TestCase): None, None, None, + None, + None, ) config = MockStorageConfig( extra_config={"dram_capacity_bytes": 1024, "ssd_enabled": False}