diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 5d05fa666..d17ac17b2 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -628,6 +628,26 @@ class MambaPoolHost(HostKVCache): element_size_list.append(conv_element_sizes[j]) return ptr_list, element_size_list + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + if self.layout not in ["page_first", "page_first_direct"]: + return False + temporal_stride = ( + self.num_mamba_layers + * self.temporal_state_elem_size + * self.temporal_dtype.itemsize + ) + if self.temporal_buffer.data_ptr() % page_size_bytes != 0: + return False + if temporal_stride % page_size_bytes != 0: + return False + for buf, elem_size in zip(self.conv_buffer, self.conv_state_elem_sizes): + conv_stride = self.num_mamba_layers * elem_size * self.conv_dtype.itemsize + if buf.data_ptr() % page_size_bytes != 0: + return False + if conv_stride % page_size_bytes != 0: + return False + return True + # ---- V4 Compressed KV Host Pools ---- @@ -1091,6 +1111,15 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache): return ptr_list, [page_bytes] * len(ptr_list) raise ValueError(f"Unsupported layout: {self.layout}") + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + if self.layout not in ["page_first", "page_first_direct"]: + return False + page_bytes = self.layer_num * self.item_bytes * self.dtype.itemsize + return ( + self.kv_buffer.data_ptr() % page_size_bytes == 0 + and page_bytes % page_size_bytes == 0 + ) + class DeepSeekV4StateHostPool(HostKVCache): """Host pool for V4 CompressStatePool page rows.""" @@ -1460,6 +1489,15 @@ class DeepSeekV4StateHostPool(HostKVCache): return ptr_list, [page_bytes] * len(ptr_list) raise ValueError(f"Unsupported layout: {self.layout}") + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + if self.layout not in ["page_first", "page_first_direct"]: + return False + page_bytes = self.layer_num * self.state_page_bytes * self.dtype.itemsize + return ( + self.kv_buffer.data_ptr() % page_size_bytes == 0 + and page_bytes % page_size_bytes == 0 + ) + @dataclass class PoolEntry: @@ -1535,6 +1573,9 @@ class HostPoolGroup: def get_page_buffer_meta(self, indices): return self.anchor_entry.host_pool.get_page_buffer_meta(indices) + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + return self.anchor_entry.host_pool.is_stride_page_aligned(page_size_bytes) + def clear(self) -> None: for entry in self.entries: entry.host_pool.clear() @@ -1997,3 +2038,14 @@ class DSAIndexerPoolHost(HostKVCache): page_index = int(indices[i]) // self.page_size ptr_list.append(base_ptr + page_index * page_stride_bytes) return ptr_list, [page_stride_bytes] * len(ptr_list) + + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + if self.layout not in ["page_first", "page_first_direct"]: + return False + page_stride_bytes = ( + self.layer_num * self.indexer_page_stride_size * self.indexer_dtype.itemsize + ) + return ( + self.index_k_with_scale_buffer.data_ptr() % page_size_bytes == 0 + and page_stride_bytes % page_size_bytes == 0 + ) diff --git a/python/sglang/srt/mem_cache/storage/nixl/README.md b/python/sglang/srt/mem_cache/storage/nixl/README.md index f55b2b18d..0d797fa5e 100644 --- a/python/sglang/srt/mem_cache/storage/nixl/README.md +++ b/python/sglang/srt/mem_cache/storage/nixl/README.md @@ -192,6 +192,72 @@ This method is convenient for testing / experimenting. For production or multi-p Also note that the flat inline config form is interpreted as plugin-specific parameters for the selected plugin. +### 4. Validated Hybrid-Model Example + +The following setup was validated against a hybrid Mamba model with HiCache enabled: + +- model: `Qwen/Qwen3.5-9B` +- storage backend: `nixl` +- NIXL plugin: `POSIX` +- HiCache layout: `page_first_direct` +- model type: hybrid attention + Mamba sidecar cache (`KV + MAMBA`) + +Important details from this validation: + +- Use a real `.toml` file path with `--hicache-storage-backend-extra-config`. +- For this validated path, the storage directory was provided through `SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR`. +- Use `--mamba-scheduler-strategy extra_buffer` to support page sizes larger than 1. + +Example TOML file: + +```toml +[plugin.posix] +active = true +``` + +Example serve command for a hybrid model: + +```bash +export SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR=/tmp/sglang_nixl_e2e_storage + +~/ve_sgl_dev/bin/sglang serve \ + --model-path /workspace/LLM_models/Qwen3.5-9B \ + --served-model-name Qwen/Qwen3.5-9B \ + --host 127.0.0.1 \ + --tp 2 \ + --reasoning-parser qwen3 \ + --attention-backend triton \ + --enable-hierarchical-cache \ + --hicache-ratio 2 \ + --hicache-io-backend direct \ + --hicache-mem-layout page_first_direct \ + --hicache-storage-prefetch-policy wait_complete \ + --page-size 256 \ + --log-level info \ + --disable-cuda-graph \ + --hicache-storage-backend nixl \ + --hicache-storage-backend-extra-config @/tmp/nixl.config.toml \ + --mamba-scheduler-strategy extra_buffer +``` + +Expected behavior for this validated setup: + +- the server starts with `Attached hybrid Mamba pool stack to HiMambaRadixCache: pools=KV + MAMBA` +- NIXL logs show `Backend POSIX was instantiated` +- the server logs `HiCacheNixl: registered hybrid host pool mamba zero_copy=...` +- the storage directory contains KV files plus Mamba sidecar files such as `..._0_2_mamba_temporal` and `..._0_2_mamba_conv_0` +- after restarting the server against the same storage directory, a repeated long prompt shows large `cached_tokens` in the response metadata + +Minimal end-to-end validation flow: + +1. Start the server with the TOML file shown above. +2. Send a long prompt once to populate storage. +3. Restart the server against the same `SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR`. +4. Send the same long prompt again and confirm that `meta_info.cached_tokens` is high. + +A reusable local validation script is available at `~/TestEnv/nixl_hicache_hybrid_e2e.py`; it starts this server, sends a long request, and checks both NIXL backend selection and Mamba sidecar storage files. + + ## Running Unit Tests diff --git a/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py b/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py index 89d3f8bf4..f1437c57e 100644 --- a/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py +++ b/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py @@ -2,6 +2,7 @@ import logging import os import time import uuid +from dataclasses import dataclass from typing import Any, List, Optional import torch @@ -12,6 +13,10 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorage, HiCacheStorageConfig, HiCacheStorageExtraInfo, + PoolHitPolicy, + PoolName, + PoolTransfer, + PoolTransferResult, ) from sglang.srt.mem_cache.mmap_allocator import alloc_mmap from sglang.srt.mem_cache.pool_host import HostKVCache @@ -52,6 +57,15 @@ def _parse_storage_dirs(raw: Optional[str]) -> List[str]: return ordered +@dataclass +class _HybridPoolContext: + host_pool: HostKVCache + is_zero_copy: bool + bounce_set: Optional[torch.Tensor] = None + bounce_get: Optional[torch.Tensor] = None + bounce_page_bytes: int = 0 + + class HiCacheNixl(HiCacheStorage): """HiCacheNixl provides high-performance storage using NIXL plugins.""" @@ -142,6 +156,9 @@ class HiCacheNixl(HiCacheStorage): self._bounce_set: Optional[torch.Tensor] = None self._bounce_get: Optional[torch.Tensor] = None self._bounce_page_bytes: Optional[int] = None + self._logical_anchor = False + self._hybrid_pool_ctx: dict[PoolName, _HybridPoolContext] = {} + self.registered_pools: dict[PoolName, HostKVCache] = {} cleanup_dirs = ( self.file_manager.iter_all_base_dirs() if self.file_manager is not None @@ -168,12 +185,56 @@ class HiCacheNixl(HiCacheStorage): def _get_suffixed_key(self, key: str) -> str: return key + self.config_suffix + def _get_component_key( + self, key: str, component_name: Optional[PoolName] = None + ) -> str: + if component_name in (None, PoolName.KV): + return self._get_suffixed_key(key) + return f"{self._get_suffixed_key(key)}_{component_name}" + + def _get_component_keys( + self, keys: List[str], pool_name: Optional[PoolName] = None + ) -> List[str]: + return [self._get_component_key(key, pool_name) for key in keys] + + def _get_hybrid_component_keys( + self, keys: List[str], pool_name: PoolName, key_multiplier: int + ) -> List[str]: + if key_multiplier == 1: + return self._get_component_keys(keys, pool_name) + + if pool_name == PoolName.MAMBA: + suffixes = [f"_{pool_name}_temporal"] + [ + f"_{pool_name}_conv_{i}" for i in range(key_multiplier - 1) + ] + elif key_multiplier == 2: + suffixes = [f"_{pool_name}_k", f"_{pool_name}_v"] + else: + suffixes = [f"_{pool_name}_{i}" for i in range(key_multiplier)] + + return [ + f"{self._get_suffixed_key(key)}{suffix}" + for key in keys + for suffix in suffixes + ] + def _create_query_tuple(self, key: str) -> tuple: """Build the NIXL query_memory tuple for a single key.""" if self.backend_selector.mem_type == "FILE": return (0, 0, 0, self.file_manager.get_file_path(key)) return (0, 0, 0, key) + def _query_keys_exist(self, keys: List[str]) -> List[bool]: + if not keys: + return [] + tuples = [self._create_query_tuple(key) for key in keys] + query_res = self.agent.query_memory( + tuples, + self.backend_selector.backend_name, + mem_type=self.backend_selector.mem_type, + ) + return [res is not None for res in query_res] + def _xfer_and_wait( self, host_descs: Any, @@ -273,6 +334,7 @@ class HiCacheNixl(HiCacheStorage): def register_mem_pool_host(self, mem_pool_host: HostKVCache): super().register_mem_pool_host(mem_pool_host) + self._logical_anchor = False # enable zero-copy automatically if mem layout is page_first or page_first_direct self.is_zero_copy = self.mem_pool_host.layout in [ @@ -280,6 +342,30 @@ class HiCacheNixl(HiCacheStorage): "page_first_direct", ] + kv = getattr(mem_pool_host, "kv_buffer", None) + if kv is None: + # DeepSeek V4 uses a LogicalHostPool as the KV anchor. It has no + # actual KV bytes; component pools carry the data through v2 APIs. + # Still write a small marker object per page so batch_exists_v2 can + # use the anchor key to gate sidecar lookups. + self.is_zero_copy = False + self._logical_anchor = True + marker_numel = 4096 if self.needs_page_alignment else 1 + pin_memory = bool(getattr(mem_pool_host, "pin_memory", False)) + self._bounce_page_bytes = marker_numel + self._bounce_set = self._alloc_registered( + marker_numel, torch.uint8, pin_memory, "logical_anchor_set" + ) + self._bounce_get = self._alloc_registered( + marker_numel, torch.uint8, pin_memory, "logical_anchor_get" + ) + self._bounce_set.fill_(1) + logger.info( + "HiCacheNixl: registered logical anchor pool with %d-byte markers", + self._bounce_page_bytes, + ) + return + if self.needs_page_alignment and self.is_zero_copy: # Check that the kv_buffer base AND per-page strides are multiples of # the OS page size so every pointer passed to NIXL (base + p * stride) @@ -297,7 +383,6 @@ class HiCacheNixl(HiCacheStorage): self.is_zero_copy = False if self.is_zero_copy: - kv = mem_pool_host.kv_buffer self._pre_register_host( kv.data_ptr(), kv.numel() * kv.element_size(), "kv_buffer" ) @@ -322,6 +407,179 @@ class HiCacheNixl(HiCacheStorage): f"layout={mem_pool_host.layout} zero_copy={self.is_zero_copy}" ) + def register_mem_host_pool_v2(self, host_pool: HostKVCache, host_pool_name): + if host_pool_name == PoolName.KV: + return + super().register_mem_host_pool_v2(host_pool, host_pool_name) + + is_zero_copy = self._hybrid_pool_supports_zero_copy(host_pool, host_pool_name) + if is_zero_copy: + for i, buf in enumerate(host_pool.get_hybrid_pool_buffer()): + self._pre_register_host( + buf.data_ptr(), + buf.numel() * buf.element_size(), + f"{host_pool_name}_buffer_{i}", + ) + self._hybrid_pool_ctx[host_pool_name] = _HybridPoolContext( + host_pool=host_pool, is_zero_copy=True + ) + else: + sample = host_pool.get_dummy_flat_data_page() + page_numel = sample.numel() + page_bytes = page_numel * sample.element_size() + del sample + + pin_memory = bool(getattr(host_pool, "pin_memory", False)) + bounce_set = self._alloc_registered( + page_numel, host_pool.dtype, pin_memory, f"{host_pool_name}_bounce_set" + ) + bounce_get = self._alloc_registered( + page_numel, host_pool.dtype, pin_memory, f"{host_pool_name}_bounce_get" + ) + self._hybrid_pool_ctx[host_pool_name] = _HybridPoolContext( + host_pool=host_pool, + is_zero_copy=False, + bounce_set=bounce_set, + bounce_get=bounce_get, + bounce_page_bytes=page_bytes, + ) + + logger.info( + "HiCacheNixl: registered hybrid host pool %s zero_copy=%s", + host_pool_name, + is_zero_copy, + ) + + def _hybrid_pool_supports_zero_copy( + self, host_pool: HostKVCache, host_pool_name: PoolName + ) -> bool: + if not ( + hasattr(host_pool, "get_page_buffer_meta") + and hasattr(host_pool, "get_hybrid_pool_buffer") + ): + return False + buffers = host_pool.get_hybrid_pool_buffer() + if not buffers: + return False + if self.needs_page_alignment and not host_pool.is_stride_page_aligned(4096): + logger.warning( + "HiCacheNixl: O_DIRECT is active but hybrid pool %s is not " + "OS-page-aligned. Falling back to bounce buffers.", + host_pool_name, + ) + return False + return True + + def _get_bounce_slot_buffers( + self, buf: torch.Tensor, page_bytes: int, page_num: int + ) -> List[tuple]: + base = buf.data_ptr() + return [(base + i * page_bytes, page_bytes) for i in range(page_num)] + + def _get_hybrid_key_multiplier( + self, pool_name: PoolName, host_pool: HostKVCache + ) -> int: + if pool_name == PoolName.MAMBA: + return 1 + len(getattr(host_pool, "conv_buffer", []) or []) + if hasattr(host_pool, "v_buffer"): + return 2 + return 1 + + def _get_hybrid_zero_copy_buffers( + self, transfer: PoolTransfer, ctx: _HybridPoolContext + ) -> tuple[List[str], List[tuple], int]: + """Build NIXL keys and memory descriptors for zero-copy hybrid transfers. + + The host pool returns one or more physical buffers per logical cache page + depending on the pool type, for example K/V buffers for SWA or temporal + plus convolution buffers for Mamba. This helper expands each logical page + key into component-level storage keys, validates that the expanded keys + match the host-pool metadata, and returns `(key_strs, host_buffers, + key_multiplier)`. + """ + ptr_list, size_list = ctx.host_pool.get_page_buffer_meta(transfer.host_indices) + page_num = len(transfer.keys or []) + if page_num == 0 or len(ptr_list) % page_num != 0: + logger.error( + "HiCacheNixl: hybrid pool %s metadata mismatch: pages=%s ptrs=%s", + transfer.name, + page_num, + len(ptr_list), + ) + return [], [], 0 + key_multiplier = len(ptr_list) // page_num + key_strs = self._get_hybrid_component_keys( + transfer.keys or [], transfer.name, key_multiplier + ) + if len(key_strs) != len(ptr_list): + logger.error( + "HiCacheNixl: hybrid pool %s key/meta mismatch: keys=%s ptrs=%s", + transfer.name, + len(key_strs), + len(ptr_list), + ) + return [], [], 0 + return key_strs, list(zip(ptr_list, size_list)), key_multiplier + + def _prepare_pool_transfer( + self, transfer: PoolTransfer, for_write: bool + ) -> tuple[Optional[HostKVCache], List[str], List[tuple], List[int], int]: + ctx = self._hybrid_pool_ctx.get(transfer.name) + if ctx is None: + logger.error("Host pool %s is not registered in HiCacheNixl", transfer.name) + return None, [], [], [], 0 + + host_pool = ctx.host_pool + keys = transfer.keys or [] + host_indices = transfer.host_indices + page_size = getattr(host_pool, "page_size", 1) or 1 + expected = len(keys) * page_size + if host_indices is None or host_indices.numel() != expected: + logger.error( + "Pool %s indices length mismatch: expected %s, got %s", + transfer.name, + expected, + host_indices.numel() if host_indices is not None else 0, + ) + return host_pool, [], [], [], 0 + + if ctx.is_zero_copy: + key_strs, host_buffers, key_multiplier = self._get_hybrid_zero_copy_buffers( + transfer, ctx + ) + page_offsets = [ + host_indices[i * page_size].item() for i in range(len(keys)) + ] + return host_pool, key_strs, host_buffers, page_offsets, key_multiplier + + if len(keys) > STORAGE_BATCH_SIZE: + logger.error( + "HiCacheNixl: hybrid pool %s batch size %s exceeds bounce buffer capacity %s", + transfer.name, + len(keys), + STORAGE_BATCH_SIZE, + ) + return host_pool, [], [], [], 0 + + page_offsets = [host_indices[i * page_size].item() for i in range(len(keys))] + bounce = ctx.bounce_set if for_write else ctx.bounce_get + if bounce is None: + logger.error( + "Hybrid pool %s bounce buffer is not registered", transfer.name + ) + return host_pool, [], [], [], 0 + + if for_write: + for i, page_offset in enumerate(page_offsets): + src = host_pool.get_data_page(page_offset, flat=True) + bounce[i].copy_(src) + + host_buffers = self._get_bounce_slot_buffers( + bounce, ctx.bounce_page_bytes, len(page_offsets) + ) + key_strs = self._get_component_keys(keys, transfer.name) + return host_pool, key_strs, host_buffers, page_offsets, 1 + def _alloc_registered( self, page_numel: int, @@ -366,6 +624,7 @@ class HiCacheNixl(HiCacheStorage): self._bounce_set = None self._bounce_get = None self._bounce_page_bytes = None + self._hybrid_pool_ctx.clear() def __del__(self): try: @@ -391,18 +650,12 @@ class HiCacheNixl(HiCacheStorage): key_list = [self._get_suffixed_key(key) for key in keys] key_denominator = 1 - tuples = [self._create_query_tuple(key) for key in key_list] + exists_results = self._query_keys_exist(key_list) - query_res = self.agent.query_memory( - tuples, - self.backend_selector.backend_name, - mem_type=self.backend_selector.mem_type, - ) - - for i in range(len(query_res)): - if query_res[i] is None: + for i, exists in enumerate(exists_results): + if not exists: return i // key_denominator - return len(query_res) // key_denominator + return len(exists_results) // key_denominator def _get_key_list_from_meta(self, keys: List[str]) -> List[str]: # Each key maps to a `_k` entry, plus a `_v` entry on non-MLA models @@ -477,11 +730,14 @@ class HiCacheNixl(HiCacheStorage): bounce = self._bounce_set if op == "set" else self._bounce_get if op == "set": - for i in range(page_num): - src = self.mem_pool_host.get_data_page( - host_indices[i * page_size], flat=True - ) - bounce[i].copy_(src) + if self._logical_anchor: + bounce[:page_num].fill_(1) + else: + for i in range(page_num): + src = self.mem_pool_host.get_data_page( + host_indices[i * page_size], flat=True + ) + bounce[i].copy_(src) host_buffers = self._bounce_slot_buffers(bounce, page_num) key_list = [self._get_suffixed_key(key) for key in keys] @@ -531,6 +787,9 @@ class HiCacheNixl(HiCacheStorage): return results return [(results[2 * i] and results[2 * i + 1]) for i in range(page_num)] + if self._logical_anchor: + return results + # non zero copy: copy data from the get-side bounce buffer to mem_pool_host for i in range(page_num): if not results[i]: @@ -620,3 +879,135 @@ class HiCacheNixl(HiCacheStorage): ) return results + + def batch_exists_v2( + self, + keys: List[str], + pool_transfers: Optional[List[PoolTransfer]] = None, + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> PoolTransferResult: + 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 + if transfer.name not in self.registered_pools: + final_pages = 0 + break + + ctx = self._hybrid_pool_ctx.get(transfer.name) + if ctx is None: + final_pages = 0 + break + key_multiplier = ( + self._get_hybrid_key_multiplier(transfer.name, ctx.host_pool) + if ctx.is_zero_copy + else 1 + ) + component_keys = self._get_hybrid_component_keys( + keys[:kv_pages], transfer.name, key_multiplier + ) + exists_results = self._query_keys_exist(component_keys) + page_exists = self._page_results(exists_results, key_multiplier) + + boundary = 0 + if transfer.hit_policy == PoolHitPolicy.ALL_PAGES: + try: + boundary = page_exists.index(False) + except ValueError: + boundary = kv_pages + elif transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES: + trailing = max(1, len(transfer.keys) if transfer.keys else 1) + for prefix_len in range(kv_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) + + @staticmethod + def _page_results(results: List[bool], key_multiplier: int) -> List[bool]: + if key_multiplier <= 1: + return results + return [ + all(results[i : i + key_multiplier]) + for i in range(0, len(results), key_multiplier) + ] + + def batch_get_v2( + self, + transfers: List[PoolTransfer], + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> dict[str, List[bool]]: + results: dict[str, List[bool]] = {} + for transfer in transfers: + host_pool, key_strs, host_buffers, page_offsets, key_multiplier = ( + self._prepare_pool_transfer(transfer, for_write=False) + ) + if host_pool is None or not key_strs: + results[transfer.name] = [False] * len(transfer.keys or []) + continue + + start_time = time.perf_counter() + transfer_results = self._batch_xfer( + key_strs, key_strs, host_buffers, "READ" + ) + elapsed_ms = (time.perf_counter() - start_time) * 1000 + self._log_xfer_stats( + f"batch_get_v2[{transfer.name}]", + len(transfer.keys or []), + transfer.host_indices, + [size for _, size in host_buffers], + elapsed_ms, + ) + ctx = self._hybrid_pool_ctx[transfer.name] + page_results = self._page_results(transfer_results, key_multiplier) + if not ctx.is_zero_copy: + for ok, page_offset, data_page in zip( + page_results, page_offsets, ctx.bounce_get + ): + if not ok: + break + host_pool.set_from_flat_data_page(page_offset, data_page) + results[transfer.name] = page_results + return results + + def batch_set_v2( + self, + transfers: List[PoolTransfer], + extra_info: Optional[HiCacheStorageExtraInfo] = None, + ) -> dict[str, List[bool]]: + results: dict[str, List[bool]] = {} + for transfer in transfers: + _, key_strs, host_buffers, _, key_multiplier = self._prepare_pool_transfer( + transfer, for_write=True + ) + if not key_strs: + results[transfer.name] = [False] * len(transfer.keys or []) + continue + + start_time = time.perf_counter() + transfer_results = self._batch_xfer( + key_strs, key_strs, host_buffers, "WRITE" + ) + elapsed_ms = (time.perf_counter() - start_time) * 1000 + self._log_xfer_stats( + f"batch_set_v2[{transfer.name}]", + len(transfer.keys or []), + transfer.host_indices, + [size for _, size in host_buffers], + elapsed_ms, + ) + results[transfer.name] = self._page_results( + transfer_results, key_multiplier + ) + return results diff --git a/test/registered/unit/mem_cache/test_hicache_nixl_storage.py b/test/registered/unit/mem_cache/test_hicache_nixl_storage.py index 63e530fc6..3452f111e 100644 --- a/test/registered/unit/mem_cache/test_hicache_nixl_storage.py +++ b/test/registered/unit/mem_cache/test_hicache_nixl_storage.py @@ -14,7 +14,11 @@ import unittest import torch -from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig +from sglang.srt.mem_cache.hicache_storage import ( + HiCacheStorageConfig, + PoolName, + PoolTransfer, +) from sglang.srt.mem_cache.storage.nixl.hicache_nixl import HiCacheNixl from sglang.test.test_utils import CustomTestCase @@ -22,6 +26,56 @@ from sglang.test.test_utils import CustomTestCase STRESS_ENABLED = bool(os.environ.get("SGLANG_RUN_NIXL_STRESS")) +class MockHybridPool: + def __init__( + self, + num_pages: int = 4, + page_size: int = 1, + component_bytes: int = 8, + expose_zero_copy: bool = True, + ): + self.page_size = page_size + self.dtype = torch.uint8 + self.device = "cpu" + self.pin_memory = False + self.temporal_buffer = torch.zeros( + (num_pages * page_size, component_bytes), dtype=self.dtype + ) + self.conv_buffer = [ + torch.zeros((num_pages * page_size, component_bytes), dtype=self.dtype) + ] + if expose_zero_copy: + self.get_hybrid_pool_buffer = self._get_hybrid_pool_buffer + + def _get_hybrid_pool_buffer(self): + return [self.temporal_buffer, *self.conv_buffer] + + def get_page_buffer_meta(self, indices): + ptr_list = [] + size_list = [] + for index in indices.tolist(): + ptr_list.append(self.temporal_buffer[index].data_ptr()) + size_list.append(self.temporal_buffer[index].numel()) + ptr_list.append(self.conv_buffer[0][index].data_ptr()) + size_list.append(self.conv_buffer[0][index].numel()) + return ptr_list, size_list + + def get_dummy_flat_data_page(self): + return torch.zeros(self.temporal_buffer.shape[1] * 2, dtype=self.dtype) + + def get_data_page(self, index, flat=True): + data = torch.cat([self.temporal_buffer[index], self.conv_buffer[0][index]]) + return data.flatten() if flat else data + + def set_from_flat_data_page(self, index, data_page): + split = self.temporal_buffer.shape[1] + self.temporal_buffer[index].copy_(data_page[:split]) + self.conv_buffer[0][index].copy_(data_page[split:]) + + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + return True + + class MockMemPoolHost: """Minimal MHA-style HostKVCache stand-in supporting the v1 paths. @@ -511,6 +565,128 @@ class TestNixlUnified(CustomTestCase): self.assertEqual(self.hicache.batch_exists(["key1", "key2"]), 1) + def test_register_mem_host_pool_v2_uses_zero_copy_when_supported(self): + pool = MockHybridPool(expose_zero_copy=True) + self.hicache.register_mem_host_pool_v2(pool, PoolName.MAMBA) + + ctx = self.hicache._hybrid_pool_ctx[PoolName.MAMBA] + self.assertTrue(ctx.is_zero_copy) + self.assertIs(ctx.host_pool, pool) + + def test_register_mem_host_pool_v2_uses_persistent_bounce_otherwise(self): + pool = MockHybridPool(expose_zero_copy=False) + self.hicache.register_mem_host_pool_v2(pool, PoolName.MAMBA) + + ctx = self.hicache._hybrid_pool_ctx[PoolName.MAMBA] + self.assertFalse(ctx.is_zero_copy) + self.assertIsNotNone(ctx.bounce_set) + self.assertIsNotNone(ctx.bounce_get) + self.assertEqual(ctx.bounce_page_bytes, pool.get_dummy_flat_data_page().numel()) + + def test_batch_set_v2_expands_zero_copy_mamba_component_keys(self): + pool = MockHybridPool(expose_zero_copy=True) + self.hicache.register_mem_host_pool_v2(pool, PoolName.MAMBA) + + captured = {} + + def fake_batch_xfer(keys, key_strs, host_buffers, direction): + captured["keys"] = key_strs + captured["host_buffers"] = host_buffers + captured["direction"] = direction + return [True] * len(key_strs) + + self.hicache._batch_xfer = fake_batch_xfer + results = self.hicache.batch_set_v2( + [ + PoolTransfer( + name=PoolName.MAMBA, + keys=["p0", "p1"], + host_indices=torch.tensor([0, 1], dtype=torch.int64), + ) + ] + ) + + self.assertEqual(results[PoolName.MAMBA], [True, True]) + self.assertEqual( + captured["keys"], + [ + self.hicache._get_suffixed_key("p0") + "_mamba_temporal", + self.hicache._get_suffixed_key("p0") + "_mamba_conv_0", + self.hicache._get_suffixed_key("p1") + "_mamba_temporal", + self.hicache._get_suffixed_key("p1") + "_mamba_conv_0", + ], + ) + self.assertEqual(len(captured["host_buffers"]), 4) + self.assertEqual(captured["direction"], "WRITE") + + def test_batch_get_v2_uses_bounce_buffer_for_non_zero_copy_pool(self): + pool = MockHybridPool(expose_zero_copy=False) + self.hicache.register_mem_host_pool_v2(pool, PoolName.MAMBA) + + def fake_batch_xfer(keys, key_strs, host_buffers, direction): + ctx = self.hicache._hybrid_pool_ctx[PoolName.MAMBA] + ctx.bounce_get[0].fill_(3) + return [True] * len(key_strs) + + self.hicache._batch_xfer = fake_batch_xfer + results = self.hicache.batch_get_v2( + [ + PoolTransfer( + name=PoolName.MAMBA, + keys=["p0"], + host_indices=torch.tensor([0], dtype=torch.int64), + ) + ] + ) + + self.assertEqual(results[PoolName.MAMBA], [True]) + self.assertTrue(torch.all(pool.get_data_page(0) == 3)) + + def test_batch_set_get_v2_distinguishes_same_key_by_pool_name(self): + mamba_pool = MockHybridPool(expose_zero_copy=False) + swa_pool = MockHybridPool(expose_zero_copy=False) + self.hicache.register_mem_host_pool_v2(mamba_pool, PoolName.MAMBA) + self.hicache.register_mem_host_pool_v2(swa_pool, PoolName.SWA) + + mamba_pool.temporal_buffer[0].fill_(11) + mamba_pool.conv_buffer[0][0].fill_(12) + swa_pool.temporal_buffer[0].fill_(21) + swa_pool.conv_buffer[0][0].fill_(22) + expected_mamba = mamba_pool.get_data_page(0).clone() + expected_swa = swa_pool.get_data_page(0).clone() + + key = "shared_key" + host_indices = torch.tensor([0], dtype=torch.int64) + set_results = self.hicache.batch_set_v2( + [ + PoolTransfer( + name=PoolName.MAMBA, keys=[key], host_indices=host_indices + ), + PoolTransfer(name=PoolName.SWA, keys=[key], host_indices=host_indices), + ] + ) + self.assertEqual(set_results[PoolName.MAMBA], [True]) + self.assertEqual(set_results[PoolName.SWA], [True]) + + mamba_pool.temporal_buffer.zero_() + mamba_pool.conv_buffer[0].zero_() + swa_pool.temporal_buffer.zero_() + swa_pool.conv_buffer[0].zero_() + + get_results = self.hicache.batch_get_v2( + [ + PoolTransfer( + name=PoolName.MAMBA, keys=[key], host_indices=host_indices + ), + PoolTransfer(name=PoolName.SWA, keys=[key], host_indices=host_indices), + ] + ) + + self.assertEqual(get_results[PoolName.MAMBA], [True]) + self.assertEqual(get_results[PoolName.SWA], [True]) + self.assertTrue(torch.equal(mamba_pool.get_data_page(0), expected_mamba)) + self.assertTrue(torch.equal(swa_pool.get_data_page(0), expected_swa)) + @unittest.skipUnless(hasattr(os, "O_DIRECT"), "O_DIRECT not available on this platform") class TestNixlDirectIO(CustomTestCase):