From a34f81251f78bab5b818144649a45c6bd51a2665 Mon Sep 17 00:00:00 2001 From: AMD-yanfeiwang Date: Thu, 13 Aug 2026 15:40:07 +0800 Subject: [PATCH] =?UTF-8?q?fix(hicache/umbp):=20support=20DeepSeek-V4=20hy?= =?UTF-8?q?brid=20HostPoolGroup=20(multi-po=E2=80=A6=20(#30762)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: Zhangheng --- .../sglang/srt/mem_cache/memory_pool_host.py | 4 + .../srt/mem_cache/storage/umbp/umbp_store.py | 296 +++++++++++++++++- .../test_hicache_storage_umbp_backend.py | 214 +++++++++++++ ...test_hicache_staged_write_back_dispatch.py | 16 + .../unit/mem_cache/test_umbp_store.py | 205 +++++++++++- test/run_suite.py | 1 + 6 files changed, 706 insertions(+), 30 deletions(-) create mode 100644 test/registered/hicache/test_hicache_storage_umbp_backend.py diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 4555ced64..1835c248a 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -671,6 +671,10 @@ class LogicalHostPool: def clear(self): self.free_slots = torch.arange(self.size, dtype=torch.int64) + def destroy(self) -> None: + """Logical anchors own no backing buffers or registrations to release.""" + return None + def available_size(self): return len(self.free_slots) 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 5910b75ea..f64f096f3 100644 --- a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py @@ -18,6 +18,10 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorage, HiCacheStorageConfig, HiCacheStorageExtraInfo, + PoolHitPolicy, + PoolName, + PoolTransfer, + PoolTransferResult, ) from sglang.srt.mem_cache.memory_pool_host import HostKVCache @@ -573,6 +577,15 @@ class UMBPStore(HiCacheStorage): # .default_dram_page_size (2 MiB by default). The # partial-tail safety net in PoolClient handles any # size mismatch. + # + # Logical-anchor host pools (the DeepSeek-V4 HiCache HostPoolGroup + # whose KV anchor is a LogicalHostPool that owns only page indices + # and no physical KV tensor) return None from get_page_buffer_meta() + # by design — the real per-page byte sizes live in the v2 side pools + # (SWA / compressed KV / indexer / state), which each carry their own + # dimensions. There is no single page size that fits all of them, so + # we leave dram_page_size at 0 and let the mori master use its + # default with the PoolClient partial-tail safety net. page_byte_size = None if "dram_page_size" in extra: page_byte_size = int(extra["dram_page_size"]) @@ -584,14 +597,17 @@ class UMBPStore(HiCacheStorage): # 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) + meta = 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) + meta = mem_pool_host.get_split_heads_page_buffer_meta(dummy, sf) else: - _, esz = mem_pool_host.get_page_buffer_meta(dummy) + 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. + esz = meta[1] if meta else None page_byte_size = int(esz[0]) if esz else 0 if ( @@ -787,6 +803,12 @@ 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. + self.registered_pools: dict = {} + self._kv_anchor_is_logical = False + self.client = UMBPClient(cfg) if mem_pool_host is not None: self.register_mem_pool_host(mem_pool_host) @@ -867,23 +889,44 @@ 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(). + self._kv_anchor_is_logical = self.mem_pool_host.kv_buffer is None + + self._zero_copy_registered = False + 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. - self._zero_copy_registered = False + if self._register_host_buffer_for_zero_copy(mem_pool_host): + self._zero_copy_registered = True + + 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). + """ if self.client is None: - return + return False try: is_distributed = bool(self.client.is_distributed()) except Exception: is_distributed = False if not is_distributed: - return + return False if not hasattr(self.client, "register_memory"): - return + return False if getattr(self, "_disable_zero_copy_register", False): logger.info( "UMBPStore: skipping host KV buffer RDMA registration because " @@ -891,9 +934,20 @@ class UMBPStore(HiCacheStorage): "Falling back to the staging-buffer transfer path; per-transfer " "size is capped by distributed.staging_buffer_size." ) - return + 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: + return False 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 @@ -901,7 +955,7 @@ class UMBPStore(HiCacheStorage): # 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) + 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) @@ -917,20 +971,37 @@ class UMBPStore(HiCacheStorage): "distributed.staging_buffer_size.", exc, ) - return + return False 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." - ) + return True + logger.warning( + "UMBPStore: register_memory returned false; staying on staging " + "buffer fallback path." + ) + return False + + 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 + self._register_host_buffer_for_zero_copy(host_pool) # ------------------------------------------------------------------ # Key suffix generation — mirrors MooncakeStore @@ -1007,6 +1078,11 @@ class UMBPStore(HiCacheStorage): host_indices: torch.Tensor, 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) # Normalize sizes to list of per-key sizes @@ -1076,6 +1152,11 @@ class UMBPStore(HiCacheStorage): page_count = len(host_indices) // self.mem_pool_host.page_size 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) if isinstance(buffer_sizes, int): @@ -1138,6 +1219,187 @@ class UMBPStore(HiCacheStorage): hit_count = self.client.batch_exists_consecutive(query_keys) 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). + # ------------------------------------------------------------------ + 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. + """ + 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}"] + 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", + ] + else: + suffixes = [f"_{self.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 + + def batch_exists_v2( + self, + keys: List[str], + 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) + + 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( + keys[:final_pages], transfer + ) + exists = list(self.client.batch_exists(component_keys)) + if len(exists) != len(component_keys): + logger.error( + "UMBP v2 batch_exists result-size mismatch for pool %s: " + "expected=%d actual=%d; treating the storage prefix as a miss", + transfer.name, + len(component_keys), + len(exists), + ) + final_pages = 0 + break + # Collapse per-object results into per-page presence. + page_exists = [ + all(exists[i * key_multiplier : (i + 1) * key_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 + 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): + if all( + page_exists[i] + for i in range(max(0, prefix_len - trailing), prefix_len) + ): + boundary = prefix_len + break + if boundary: + hit_count[transfer.name] = boundary + final_pages = min(final_pages, boundary) + + return PoolTransferResult(final_pages, hit_count) + + def _batch_io_v2(self, transfers: List[PoolTransfer], is_set: bool) -> dict: + """Unified per-pool zero-copy I/O. Returns {pool_name: per-page bools}.""" + results: dict = {} + for transfer in transfers: + host_pool = self.registered_pools.get(transfer.name) + if host_pool is None: + raise ValueError(f"Unregistered UMBP hybrid pool: {transfer.name}") + keys = transfer.keys or [] + host_indices = transfer.host_indices + page_size = getattr(host_pool, "page_size", 1) or 1 + if not keys or host_indices is None: + results[transfer.name] = [False] * len(keys) + continue + assert len(keys) == len(host_indices) // page_size + + key_strs, key_multiplier = self._get_hybrid_page_component_keys( + keys, transfer + ) + 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) + ) + ] + + if len(io_results) != len(key_strs): + logger.error( + "UMBP v2 %s result-size mismatch for pool %s: " + "expected=%d actual=%d; treating every page as failed", + "set" if is_set else "get", + transfer.name, + len(key_strs), + len(io_results), + ) + 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]) + for i in range(len(keys)) + ] + return results + + def batch_get_v2( + self, + transfers: List[PoolTransfer], + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> dict: + return self._batch_io_v2(transfers, is_set=False) + + def batch_set_v2( + self, + transfers: List[PoolTransfer], + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> dict: + return self._batch_io_v2(transfers, is_set=True) + # ------------------------------------------------------------------ # Legacy ABC interface (required by HiCacheStorage) # ------------------------------------------------------------------ diff --git a/test/registered/hicache/test_hicache_storage_umbp_backend.py b/test/registered/hicache/test_hicache_storage_umbp_backend.py new file mode 100644 index 000000000..8145e262d --- /dev/null +++ b/test/registered/hicache/test_hicache_storage_umbp_backend.py @@ -0,0 +1,214 @@ +"""E2E test for DeepSeek-V4 HiCache storage with the UMBP backend. + +The first request writes the hybrid HostPoolGroup side pools to UMBP. After +flushing the device and host radix caches, the same prompt must be restored +from UMBP and report a storage-tier cache hit. + +Usage: + python3 -m pytest \ + test/registered/hicache/test_hicache_storage_umbp_backend.py -v +""" + +import json +import os +import subprocess +import time +import unittest + +import requests + +from sglang.srt.utils import is_hip, kill_process_tree +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_amd_ci( + est_time=3600, + suite="nightly-amd-8-gpu-mi35x-deepseek-v4-flash", + nightly=True, +) + +DEEPSEEK_V4_FLASH_FP8_MODEL_PATH = os.environ.get( + "DEEPSEEK_V4_FP8_MODEL_PATH", "sgl-project/DeepSeek-V4-Flash-FP8" +) +SERVER_LAUNCH_TIMEOUT = 3600 +PAGE_SIZE = 256 +TP_SIZE = 8 + + +@unittest.skipUnless(is_hip(), "UMBP HiCache requires ROCm.") +@unittest.skipUnless( + os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "unified_kv_triton") + == "unified_kv_triton", + "UMBP HiCache E2E only runs in the unified_kv_triton DSV4 nightly leg.", +) +class TestHiCacheStorageUMBPBackend(CustomTestCase): + """DeepSeek-V4 hybrid HostPoolGroup round trip through local UMBP L3.""" + + input_ids = list(range(4000, 5024)) + + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V4_FLASH_FP8_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = None + + try: + cls._launch_server() + except Exception: + cls._stop_server() + raise + + @classmethod + def tearDownClass(cls): + cls._stop_server() + + @classmethod + def _launch_server(cls): + storage_config = { + "dram_capacity_bytes": 1 * 1024 * 1024 * 1024, + "ssd_enabled": True, + "ssd_storage_dir": "/tmp/umbp_dsv4_local", + "ssd_capacity_bytes": 20 * 1024 * 1024 * 1024, + } + other_args = [ + "--trust-remote-code", + "--tp-size", + str(TP_SIZE), + "--attention-backend", + "dsv4", + "--kv-cache-dtype", + "fp8_e4m3", + "--page-size", + str(PAGE_SIZE), + "--chunked-prefill-size", + "8192", + "--mem-fraction-static", + "0.85", + "--disable-cuda-graph", + "--disable-shared-experts-fusion", + "--enable-hierarchical-cache", + "--hicache-ratio", + "2", + "--hicache-write-policy", + "write_through", + "--hicache-storage-prefetch-policy", + "wait_complete", + "--hicache-io-backend", + "direct", + "--hicache-mem-layout", + "page_first", + "--hicache-storage-backend", + "mori", + "--hicache-storage-backend-extra-config", + json.dumps(storage_config), + "--enable-cache-report", + "--enable-metrics", + "--swa-full-tokens-ratio", + "0.1", + "--max-total-tokens", + "20000", + "--max-running-requests", + "4", + "--watchdog-timeout", + "1200", + ] + + env = os.environ.copy() + # 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.update( + { + "SGLANG_ENABLE_DETERMINISTIC_INFERENCE": "1", + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", + "SGLANG_DSV4_FP4_EXPERTS": "0", + "SGLANG_HACK_FLASHMLA_BACKEND": "unified_kv_triton", + "SGLANG_USE_ROCM700A": "0", + "AITER_BF16_FP8_MOE_BOUND": "0", + # Correctness does not depend on pre-reserved hugepages, and + # disabling them makes the E2E portable across MI35x runners. + "SGLANG_HICACHE_HOST_HUGEPAGE": "0", + "UMBP_DRAM_USE_HUGEPAGES": "0", + } + ) + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=other_args, + env=env, + ) + + @classmethod + def _stop_server(cls): + process = getattr(cls, "process", None) + if process is None: + return + if process.poll() is None: + # Give UMBP clients a chance to close their local tiers before the + # process tree is force-killed. + process.terminate() + try: + process.wait(timeout=60) + except subprocess.TimeoutExpired: + kill_process_tree(process.pid) + cls.process = None + + def _flush_device_and_host_cache(self): + response = requests.post( + self.base_url + "/flush_cache", + params={"timeout": 60}, + timeout=90, + ) + response.raise_for_status() + + def _generate(self): + response = requests.post( + self.base_url + "/generate", + json={ + "input_ids": self.input_ids, + "sampling_params": { + "temperature": 0, + "max_new_tokens": 8, + "ignore_eos": True, + }, + }, + timeout=1200, + ) + self.assertEqual( + response.status_code, + 200, + f"Request failed: {response.status_code} - {response.text}", + ) + return response.json() + + def test_hybrid_host_pool_round_trip_from_umbp(self): + self._flush_device_and_host_cache() + + first = self._generate() + self.assertEqual(first["meta_info"]["cached_tokens"], 0) + + # Writes are asynchronous below the request path. This mirrors the + # Mooncake E2E drain before forcing the next request to use L3. + time.sleep(15) + self._flush_device_and_host_cache() + + second = self._generate() + cached_details = second["meta_info"].get("cached_tokens_details") or {} + storage_cached_tokens = int(cached_details.get("storage", 0)) + + self.assertGreaterEqual( + storage_cached_tokens, + PAGE_SIZE, + "Expected DeepSeek-V4 side-pool KV to load from UMBP storage, " + f"got {cached_details=}", + ) + self.assertEqual(cached_details.get("storage_backend"), "UMBPStore") + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index 57c0353bd..0a846b2ec 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -684,6 +684,22 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): self.assertEqual(group.layout, "page_first") self.assertTrue(group.can_use_write_back_jit) + def test_host_pool_group_destroys_logical_anchor(self): + logical_host_pool = LogicalHostPool(8, 2, layout="page_first") + group = HostPoolGroup( + [ + PoolEntry( + name=PoolName.KV, + host_pool=logical_host_pool, + device_pool=None, + layer_mapper=lambda _: 0, + is_primary_index_anchor=True, + ) + ] + ) + + self.assertIsNone(group.destroy()) + def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self): captured = {} diff --git a/test/registered/unit/mem_cache/test_umbp_store.py b/test/registered/unit/mem_cache/test_umbp_store.py index 142840a76..699cd3213 100755 --- a/test/registered/unit/mem_cache/test_umbp_store.py +++ b/test/registered/unit/mem_cache/test_umbp_store.py @@ -2,25 +2,20 @@ """Unit tests for UMBPStore with mocked HostKVCache.""" import ctypes +import importlib +import sys import tempfile import unittest from dataclasses import dataclass +from types import ModuleType, SimpleNamespace from typing import Optional -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch -from sglang.test.ci.ci_register import register_cpu_ci +import mori.umbp # noqa: F401 -register_cpu_ci(est_time=5, suite="base-a-test-cpu") +from sglang.test.ci.ci_register import register_amd_ci -# 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 +register_amd_ci(est_time=30, suite="stage-a-test-1-gpu-small-amd") @dataclass @@ -90,12 +85,39 @@ class MockHostKVCache: return bytes(ctypes.string_at(self._buffer_ptr + v_offset, self.element_size)) +class MockLogicalHostPool: + layout = "page_first" + page_size = 1 + kv_buffer = None + + +class MockHybridSidePool: + page_size = 1 + + def get_page_buffer_meta(self, indices): + return [1000 + i * 8 for i in range(len(indices))], [8] * len(indices) + + +def import_umbp_store_module(): + """Import UMBPStore without pulling GPU-only memory-pool dependencies.""" + module_name = "sglang.srt.mem_cache.storage.umbp.umbp_store" + if module_name in sys.modules: + return sys.modules[module_name] + + fake_memory_pool_host = ModuleType("sglang.srt.mem_cache.memory_pool_host") + fake_memory_pool_host.HostKVCache = object + with patch.dict( + sys.modules, + {"sglang.srt.mem_cache.memory_pool_host": fake_memory_pool_host}, + ): + return importlib.import_module(module_name) + + 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 @@ -290,5 +312,162 @@ class TestUMBPStore(unittest.TestCase): store.clear() +class TestUMBPStoreDefensiveSemantics(unittest.TestCase): + @staticmethod + def _make_v2_store(): + from sglang.srt.mem_cache.hicache_storage import PoolName + + UMBPStore = import_umbp_store_module().UMBPStore + store = UMBPStore.__new__(UMBPStore) + store.client = MagicMock() + store.client.is_distributed.return_value = False + store.registered_pools = {} + store._kv_anchor_is_logical = True + store.is_mla_backend = True + store.mla_suffix = "" + store.mha_suffix = "0" + store.register_mem_host_pool_v2(MockHybridSidePool(), PoolName.DEEPSEEK_V4_C4) + return store + + def test_constructor_preserves_logical_anchor_detection(self): + umbp_module = import_umbp_store_module() + + class FakeUMBPConfig: + def __init__(self): + self.role = None + self.dram = SimpleNamespace(capacity_bytes=0) + self.ssd = SimpleNamespace( + enabled=False, + storage_dir="/tmp", + capacity_bytes=0, + ssd_backend="file", + spdk_proxy_tenant_id=0, + spdk_proxy_tenant_quota_bytes=0, + ) + self.distributed = None + + @classmethod + def from_environment(cls): + return cls() + + class FakeUMBPClient: + def __init__(self, _config): + pass + + fake_role = SimpleNamespace( + Standalone="standalone", + SharedSSDLeader="leader", + SharedSSDFollower="follower", + ) + imported = ( + FakeUMBPClient, + FakeUMBPConfig, + fake_role, + None, + None, + None, + ) + config = MockStorageConfig( + extra_config={"dram_capacity_bytes": 1024, "ssd_enabled": False} + ) + + with patch.object(umbp_module, "_import_umbp_client", return_value=imported): + store = umbp_module.UMBPStore(config, MockLogicalHostPool()) + + self.assertTrue(store._kv_anchor_is_logical) + self.assertEqual(store.batch_set_v1(["page0"], [0]), [True]) + + def test_short_batch_exists_result_fails_closed(self): + from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer + + store = self._make_v2_store() + store.client.batch_exists.return_value = [True] + transfer = PoolTransfer( + name=PoolName.DEEPSEEK_V4_C4, + keys=["page0", "page1"], + host_indices=[0, 1], + ) + + result = store.batch_exists_v2(["page0", "page1"], [transfer]) + + self.assertEqual(result.kv_hit_pages, 0) + + def test_batch_exists_v2_narrows_queries_across_side_pools(self): + from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer + + store = self._make_v2_store() + store.register_mem_host_pool_v2(MockHybridSidePool(), PoolName.DEEPSEEK_V4_C128) + page_keys = [f"page{i}" for i in range(4)] + store.client.batch_exists.side_effect = [ + [True, True, False, True], + [True, False], + ] + transfers = [ + PoolTransfer( + name=PoolName.DEEPSEEK_V4_C4, + keys=page_keys, + host_indices=[0, 1, 2, 3], + ), + PoolTransfer( + name=PoolName.DEEPSEEK_V4_C128, + keys=page_keys, + host_indices=[0, 1, 2, 3], + ), + ] + + result = store.batch_exists_v2(page_keys, transfers) + + queried_keys = [ + invocation.args[0] + for invocation in store.client.batch_exists.call_args_list + ] + self.assertEqual( + queried_keys, + [ + [f"{key}__{PoolName.DEEPSEEK_V4_C4}" for key in page_keys], + [f"{key}__{PoolName.DEEPSEEK_V4_C128}" for key in page_keys[:2]], + ], + ) + self.assertEqual(result.kv_hit_pages, 1) + self.assertEqual( + result.extra_pool_hit_pages, + { + PoolName.KV: 4, + PoolName.DEEPSEEK_V4_C4: 2, + PoolName.DEEPSEEK_V4_C128: 1, + }, + ) + + def test_short_batch_get_result_marks_every_page_failed(self): + from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer + + store = self._make_v2_store() + store.client.batch_get_into_ptr.return_value = [True] + transfer = PoolTransfer( + name=PoolName.DEEPSEEK_V4_C4, + keys=["page0", "page1"], + host_indices=[0, 1], + ) + + result = store.batch_get_v2([transfer]) + + self.assertEqual(result[PoolName.DEEPSEEK_V4_C4], [False, False]) + + def test_short_batch_set_result_marks_every_page_failed(self): + from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer + + store = self._make_v2_store() + store.client.batch_put_from_ptr.return_value = [True] + transfer = PoolTransfer( + name=PoolName.DEEPSEEK_V4_C4, + keys=["page0", "page1"], + host_indices=[0, 1], + ) + + result = store.batch_set_v2([transfer]) + + self.assertEqual(result[PoolName.DEEPSEEK_V4_C4], [False, False]) + + if __name__ == "__main__": unittest.main() diff --git a/test/run_suite.py b/test/run_suite.py index b4fb70de2..00bf0dcf6 100644 --- a/test/run_suite.py +++ b/test/run_suite.py @@ -143,6 +143,7 @@ NIGHTLY_SUITES = { "nightly-amd-4-gpu", "nightly-amd-8-gpu", "nightly-amd-vlm", + "nightly-amd-8-gpu-mi35x-deepseek-v4-flash", # MI35x 8-GPU suite (different model configs) "nightly-amd-8-gpu-mi35x", ],