From ff1fc1fbdff315fe44b9431ca5aae00d7bd7f733 Mon Sep 17 00:00:00 2001 From: shuwenn <47200617+alphabetc1@users.noreply.github.com> Date: Sat, 20 Jun 2026 20:44:08 +0800 Subject: [PATCH] [mem_cache][5/N] refactor: extract host KV cache base layer into pool_host package (#27273) --- .../sglang/srt/managers/cache_controller.py | 2 +- .../sglang/srt/mem_cache/hicache_storage.py | 2 +- .../sglang/srt/mem_cache/memory_pool_host.py | 327 +----------------- .../srt/mem_cache/pool_host/__init__.py | 7 + python/sglang/srt/mem_cache/pool_host/base.py | 177 ++++++++++ .../sglang/srt/mem_cache/pool_host/common.py | 97 ++++++ .../srt/mem_cache/pool_host/hisparse.py | 71 ++++ .../aibrix_kvcache/aibrix_kvcache_storage.py | 2 +- .../srt/mem_cache/storage/eic/eic_storage.py | 2 +- .../mem_cache/storage/hf3fs/storage_hf3fs.py | 2 +- .../storage/mooncake_store/mooncake_store.py | 7 +- .../mem_cache/storage/nixl/hicache_nixl.py | 2 +- .../mem_cache/storage/simm/hicache_simm.py | 2 +- test/registered/jit/test_hicache.py | 4 +- .../unit/managers/test_hisparse_unit.py | 4 +- .../unit/mem_cache/test_dsa_pool_host_unit.py | 4 +- 16 files changed, 380 insertions(+), 332 deletions(-) create mode 100644 python/sglang/srt/mem_cache/pool_host/__init__.py create mode 100644 python/sglang/srt/mem_cache/pool_host/base.py create mode 100644 python/sglang/srt/mem_cache/pool_host/common.py create mode 100644 python/sglang/srt/mem_cache/pool_host/hisparse.py diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 5d552abe4..ca276bc11 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -31,7 +31,7 @@ from sglang.srt.mem_cache.hicache_storage import ( if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator - from sglang.srt.mem_cache.memory_pool_host import HostKVCache + from sglang.srt.mem_cache.pool_host import HostKVCache from sglang.srt.distributed import ( get_pipeline_model_parallel_rank, diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index 193be40d2..c59bd26ae 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -14,7 +14,7 @@ import torch from sglang.srt.environ import envs if TYPE_CHECKING: - from sglang.srt.mem_cache.memory_pool_host import HostKVCache + from sglang.srt.mem_cache.pool_host import HostKVCache logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index ed1718f86..0799a6b73 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1,11 +1,8 @@ from __future__ import annotations -import abc import logging import threading -from collections import defaultdict from dataclasses import dataclass -from functools import wraps from typing import TYPE_CHECKING, Any, Callable, Optional if TYPE_CHECKING: @@ -40,12 +37,10 @@ from sglang.jit_kernel.hicache import ( from sglang.jit_kernel.hisparse import transfer_cache_dsv4_mla from sglang.srt.mem_cache.memory_pool import ( DSATokenToKVPool, - KVCache, MambaPool, MHATokenToKVPool, MLATokenToKVPool, ) -from sglang.srt.mem_cache.mmap_allocator import alloc_mmap from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu _is_cuda = is_cuda() @@ -74,321 +69,21 @@ if _is_npu: logger = logging.getLogger(__name__) -# Host RAM to leave free when sizing HiCache pools (OS, other processes). -HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3) + +from sglang.srt.mem_cache.pool_host import HostKVCache +from sglang.srt.mem_cache.pool_host.base import ( + HICACHE_HOST_MEMORY_RESERVE_BYTES, + synchronized, +) +from sglang.srt.mem_cache.pool_host.common import ( + ALLOC_MEMORY_FUNCS, + get_allocator_from_storage, +) +from sglang.srt.mem_cache.pool_host.hisparse import HiSparseHostPoolMixin _WRITE_BACK_STAGING_PAGE_CHUNK = 64 -def synchronized(func): - @wraps(func) - def wrapper(self, *args, **kwargs): - with self.lock: - return func(self, *args, **kwargs) - - return wrapper - - -class HostTensorAllocator: - def __init__(self): - """Initialize the HostTensorAllocator.""" - self.dtype = None - self.dims = None - - def allocate(self, dims: tuple, dtype: torch.dtype, device: str) -> torch.Tensor: - assert ( - device == "cpu" - ), f"HostTensorAllocator only supports CPU allocations; got device={device!r}" - self.dtype = dtype - self.dims = dims - return alloc_mmap(dims, dtype) - - -class HiSparseHostPoolMixin: - def _round_up_to_page_size(self, size: int) -> int: - return (size + self.page_size - 1) // self.page_size * self.page_size - - def alloc_page(self, num_pages: int) -> Optional[torch.Tensor]: - return self.alloc(num_pages * self.page_size) - - def alloc_paged_token_slots( - self, - req_to_host_pool: torch.Tensor, - req_to_host_pool_allocated_len: torch.Tensor, - req_pool_idx: int, - start_pos: int, - num_tokens: int, - ) -> torch.Tensor: - """Allocate request host slots by page and return token-granular slots.""" - device = req_to_host_pool.device - if num_tokens <= 0: - return torch.empty((0,), dtype=torch.int64, device=device) - - allocated_len = int(req_to_host_pool_allocated_len[req_pool_idx]) - end_pos = start_pos + num_tokens - page_end = self._round_up_to_page_size(end_pos) - assert start_pos <= allocated_len - - if page_end > allocated_len: - num_new_pages = (page_end - allocated_len) // self.page_size - host_locs = self.alloc_page(num_new_pages) - if host_locs is None: - logger.error( - "HiSparse: host mem pool alloc failed for %d host pages " - "(req_pool_idx=%d, start_pos=%d, num_tokens=%d)", - num_new_pages, - req_pool_idx, - start_pos, - num_tokens, - ) - raise RuntimeError( - f"HiSparse host mem pool alloc failed for {num_new_pages} pages" - ) - - req_to_host_pool[req_pool_idx, allocated_len:page_end] = host_locs.to( - device=device, non_blocking=True - ) - req_to_host_pool_allocated_len[req_pool_idx] = page_end - - return req_to_host_pool[req_pool_idx, start_pos:end_pos] - - def allocated_host_indices( - self, - req_to_host_pool: torch.Tensor, - req_pool_idx: int, - allocated_len: int, - ) -> torch.Tensor: - allocated_len = int(allocated_len) - host_len = min( - self._round_up_to_page_size(allocated_len), - req_to_host_pool.shape[1], - ) - host_indices = req_to_host_pool[req_pool_idx, :host_len] - return host_indices[host_indices >= 0] - - -def get_allocator_from_storage(allocator_type): - if allocator_type == "mooncake": - try: - from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import ( - MooncakeHostTensorAllocator, - ) - - return MooncakeHostTensorAllocator() - except ImportError: - logger.warning( - "Mooncake's tensor allocator requires mooncake >= 0.3.8.post1. " - "Please upgrade Mooncake by 'pip install mooncake-transfer-engine --upgrade'. " - "Fallback to use default allocator." - ) - return HostTensorAllocator() - else: - return HostTensorAllocator() - - -def _cuda_host_register(buffer: torch.Tensor) -> None: - cudart = torch.cuda.cudart() - n_bytes = buffer.numel() * buffer.element_size() - rc = cudart.cudaHostRegister(buffer.data_ptr(), n_bytes, 0) - if int(rc) != 0: - raise RuntimeError( - f"cudaHostRegister failed (rc={int(rc)}, " - f"{cudart.cudaGetErrorString(rc)}) for ptr={buffer.data_ptr():#x} " - f"size={n_bytes}; host buffer is not pinned and device transfers " - f"may silently return stale data." - ) - - -def alloc_with_host_register( - dims: tuple, - dtype: torch.dtype, - device: str, - pin_memory: bool, - allocator: HostTensorAllocator, -) -> torch.Tensor: - """ - Allocate tensor and register host memory with cudaHostRegister. - CudaHostRegister only applies when pin_memory=True. - """ - buffer = allocator.allocate(dims, dtype=dtype, device=device) - if pin_memory: - _cuda_host_register(buffer) - return buffer - - -def alloc_with_pin_memory( - dims: tuple, - dtype: torch.dtype, - device: str, - pin_memory: bool, - allocator: None, -) -> torch.Tensor: - """ - Allocate tensor using PyTorch's built-in pin_memory flag. - """ - buffer = torch.empty(dims, dtype=dtype, device=device, pin_memory=pin_memory) - return buffer - - -ALLOC_MEMORY_FUNCS = defaultdict( - lambda: alloc_with_host_register, - { - "npu": alloc_with_pin_memory, - "musa": alloc_with_pin_memory, - }, -) - - -class HostKVCache(abc.ABC): - - def __init__( - self, - device_pool: KVCache, - host_to_device_ratio: float, - host_size: int, - page_size: int, - layout: str, - pin_memory: bool, - device: str, - allocator_type: str = "default", - ): - self.device_pool = device_pool - self.page_size = page_size - self.layout = layout - self.pin_memory = pin_memory - self.device = device - self.allocator = get_allocator_from_storage(allocator_type) - self.can_use_write_back_jit = False - - self.dtype = device_pool.store_dtype - self.size_per_token = self.get_size_per_token() - if host_size > 0: - self.size = int(host_size * 1e9 // self.size_per_token) - else: - self.size = int(device_pool.size * host_to_device_ratio) - # Align up the host memory pool size to the page size - self.page_num = self.size // self.page_size + 1 - self.size = self.page_num * self.page_size - self.start_layer = device_pool.start_layer - self.end_layer = device_pool.end_layer - - assert ( - self.size > device_pool.size - ), "The host memory should be larger than the device memory with the current protocol" - - # Verify there is enough available host memory. - host_mem = psutil.virtual_memory() - requested_bytes = self.size * self.size_per_token - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES - if requested_bytes > available_bytes: - raise ValueError( - f"Not enough host memory available. Requesting " - f"{requested_bytes / 1e9:.2f} GB but only have " - f"{available_bytes / 1e9:.2f} GB free. Please reduce the " - f"size of the hierarchical cache." - ) - else: - logger.info( - f"Allocating {requested_bytes / 1e9:.2f} GB host memory for hierarchical KV cache." - ) - - self.kv_buffer = self.init_kv_buffer() - - # A lock for synchronized operations on memory allocation and state transitions. - self.lock = threading.RLock() - self.clear() - - @abc.abstractmethod - def get_size_per_token(self): - raise NotImplementedError() - - @abc.abstractmethod - def init_kv_buffer(self): - raise NotImplementedError() - - @abc.abstractmethod - def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend - ) -> None: - """ - Load KV data from the host memory pool to the device memory pool for a specific layer. - """ - raise NotImplementedError() - - @abc.abstractmethod - def backup_from_device_all_layer( - self, device_pool, host_indices, device_indices, io_backend - ) -> None: - """ - Backup KV data from the device memory pool to the host memory pool for all layers. - """ - raise NotImplementedError() - - @abc.abstractmethod - def get_data_page(self, index, flat: bool = True) -> torch.Tensor: - """ - Get a flat data page from the host memory pool. - """ - raise NotImplementedError() - - @abc.abstractmethod - def get_dummy_flat_data_page(self) -> torch.Tensor: - """ - Get a dummy flat data page from the host memory pool. - This is used for prefetching or initializing empty pages. - """ - raise NotImplementedError() - - @abc.abstractmethod - def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None: - """ - Set a flat data page to the host memory pool. - """ - raise NotImplementedError() - - def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: - """Return True if per-page strides are multiples of *page_size_bytes*. - - Subclasses should override this with a layout-specific stride formula. - This base implementation logs a warning and returns False (safe default). - """ - logger.warning( - "%s does not implement is_stride_page_aligned(); assuming not aligned. " - "O_DIRECT with a file-based NIXL backend will fall back to copy mode for this pool.", - type(self).__name__, - ) - return False - - @synchronized - def clear(self): - # Initialize memory states and tracking structures. - self.mem_state = torch.zeros( - (self.size,), dtype=torch.uint8, device=self.device - ) - self.free_slots = torch.arange(self.size, dtype=torch.int64) - - def available_size(self): - return len(self.free_slots) - - @synchronized - def alloc(self, need_size: int) -> Optional[torch.Tensor]: - assert ( - need_size % self.page_size == 0 - ), "The requested size should be a multiple of the page size." - if need_size > self.available_size(): - return None - - select_index = self.free_slots[:need_size] - self.free_slots = self.free_slots[need_size:] - - return select_index - - @synchronized - def free(self, indices: torch.Tensor) -> int: - self.free_slots = torch.cat([self.free_slots, indices.cpu()]) - return len(indices) - - class MHATokenToKVPoolHost(HostKVCache): device_pool: MHATokenToKVPool diff --git a/python/sglang/srt/mem_cache/pool_host/__init__.py b/python/sglang/srt/mem_cache/pool_host/__init__.py new file mode 100644 index 000000000..6855269eb --- /dev/null +++ b/python/sglang/srt/mem_cache/pool_host/__init__.py @@ -0,0 +1,7 @@ +from sglang.srt.mem_cache.pool_host.base import HostKVCache +from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator + +__all__ = [ + "HostKVCache", + "HostTensorAllocator", +] diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py new file mode 100644 index 000000000..bafe065db --- /dev/null +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -0,0 +1,177 @@ +from __future__ import annotations + +import abc +import logging +import threading +from functools import wraps +from typing import Optional + +import psutil +import torch + +from sglang.srt.mem_cache.memory_pool import KVCache +from sglang.srt.mem_cache.pool_host.common import get_allocator_from_storage + +logger = logging.getLogger(__name__) + +# Host RAM to leave free when sizing HiCache pools (OS, other processes). +HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3) + + +def synchronized(func): + @wraps(func) + def wrapper(self, *args, **kwargs): + with self.lock: + return func(self, *args, **kwargs) + + return wrapper + + +class HostKVCache(abc.ABC): + + def __init__( + self, + device_pool: KVCache, + host_to_device_ratio: float, + host_size: int, + page_size: int, + layout: str, + pin_memory: bool, + device: str, + allocator_type: str = "default", + ): + self.device_pool = device_pool + self.page_size = page_size + self.layout = layout + self.pin_memory = pin_memory + self.device = device + self.allocator = get_allocator_from_storage(allocator_type) + self.can_use_write_back_jit = False + + self.dtype = device_pool.store_dtype + self.size_per_token = self.get_size_per_token() + if host_size > 0: + self.size = int(host_size * 1e9 // self.size_per_token) + else: + self.size = int(device_pool.size * host_to_device_ratio) + # Align up the host memory pool size to the page size + self.page_num = self.size // self.page_size + 1 + self.size = self.page_num * self.page_size + self.start_layer = device_pool.start_layer + self.end_layer = device_pool.end_layer + + assert ( + self.size > device_pool.size + ), "The host memory should be larger than the device memory with the current protocol" + + # Verify there is enough available host memory. + host_mem = psutil.virtual_memory() + requested_bytes = self.size * self.size_per_token + available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + if requested_bytes > available_bytes: + raise ValueError( + f"Not enough host memory available. Requesting " + f"{requested_bytes / 1e9:.2f} GB but only have " + f"{available_bytes / 1e9:.2f} GB free. Please reduce the " + f"size of the hierarchical cache." + ) + else: + logger.info( + f"Allocating {requested_bytes / 1e9:.2f} GB host memory for hierarchical KV cache." + ) + + self.kv_buffer = self.init_kv_buffer() + + # A lock for synchronized operations on memory allocation and state transitions. + self.lock = threading.RLock() + self.clear() + + @abc.abstractmethod + def get_size_per_token(self): + raise NotImplementedError() + + @abc.abstractmethod + def init_kv_buffer(self): + raise NotImplementedError() + + @abc.abstractmethod + def load_to_device_per_layer( + self, device_pool, host_indices, device_indices, layer_id, io_backend + ) -> None: + """ + Load KV data from the host memory pool to the device memory pool for a specific layer. + """ + raise NotImplementedError() + + @abc.abstractmethod + def backup_from_device_all_layer( + self, device_pool, host_indices, device_indices, io_backend + ) -> None: + """ + Backup KV data from the device memory pool to the host memory pool for all layers. + """ + raise NotImplementedError() + + @abc.abstractmethod + def get_data_page(self, index, flat: bool = True) -> torch.Tensor: + """ + Get a flat data page from the host memory pool. + """ + raise NotImplementedError() + + @abc.abstractmethod + def get_dummy_flat_data_page(self) -> torch.Tensor: + """ + Get a dummy flat data page from the host memory pool. + This is used for prefetching or initializing empty pages. + """ + raise NotImplementedError() + + @abc.abstractmethod + def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None: + """ + Set a flat data page to the host memory pool. + """ + raise NotImplementedError() + + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + """Return True if per-page strides are multiples of *page_size_bytes*. + + Subclasses should override this with a layout-specific stride formula. + This base implementation logs a warning and returns False (safe default). + """ + logger.warning( + "%s does not implement is_stride_page_aligned(); assuming not aligned. " + "O_DIRECT with a file-based NIXL backend will fall back to copy mode for this pool.", + type(self).__name__, + ) + return False + + @synchronized + def clear(self): + # Initialize memory states and tracking structures. + self.mem_state = torch.zeros( + (self.size,), dtype=torch.uint8, device=self.device + ) + self.free_slots = torch.arange(self.size, dtype=torch.int64) + + def available_size(self): + return len(self.free_slots) + + @synchronized + def alloc(self, need_size: int) -> Optional[torch.Tensor]: + assert ( + need_size % self.page_size == 0 + ), "The requested size should be a multiple of the page size." + if need_size > self.available_size(): + return None + + select_index = self.free_slots[:need_size] + self.free_slots = self.free_slots[need_size:] + + return select_index + + @synchronized + def free(self, indices: torch.Tensor) -> int: + self.free_slots = torch.cat([self.free_slots, indices.cpu()]) + return len(indices) diff --git a/python/sglang/srt/mem_cache/pool_host/common.py b/python/sglang/srt/mem_cache/pool_host/common.py new file mode 100644 index 000000000..3e90b3d24 --- /dev/null +++ b/python/sglang/srt/mem_cache/pool_host/common.py @@ -0,0 +1,97 @@ +from __future__ import annotations + +import logging +from collections import defaultdict + +import torch + +from sglang.srt.mem_cache.mmap_allocator import alloc_mmap + +logger = logging.getLogger(__name__) + + +class HostTensorAllocator: + def __init__(self): + """Initialize the HostTensorAllocator.""" + self.dtype = None + self.dims = None + + def allocate(self, dims: tuple, dtype: torch.dtype, device: str) -> torch.Tensor: + assert ( + device == "cpu" + ), f"HostTensorAllocator only supports CPU allocations; got device={device!r}" + self.dtype = dtype + self.dims = dims + return alloc_mmap(dims, dtype) + + +def get_allocator_from_storage(allocator_type): + if allocator_type == "mooncake": + try: + from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import ( + MooncakeHostTensorAllocator, + ) + + return MooncakeHostTensorAllocator() + except ImportError: + logger.warning( + "Mooncake's tensor allocator requires mooncake >= 0.3.8.post1. " + "Please upgrade Mooncake by 'pip install mooncake-transfer-engine --upgrade'. " + "Fallback to use default allocator." + ) + return HostTensorAllocator() + else: + return HostTensorAllocator() + + +def _cuda_host_register(buffer: torch.Tensor) -> None: + cudart = torch.cuda.cudart() + n_bytes = buffer.numel() * buffer.element_size() + rc = cudart.cudaHostRegister(buffer.data_ptr(), n_bytes, 0) + if int(rc) != 0: + raise RuntimeError( + f"cudaHostRegister failed (rc={int(rc)}, " + f"{cudart.cudaGetErrorString(rc)}) for ptr={buffer.data_ptr():#x} " + f"size={n_bytes}; host buffer is not pinned and device transfers " + f"may silently return stale data." + ) + + +def alloc_with_host_register( + dims: tuple, + dtype: torch.dtype, + device: str, + pin_memory: bool, + allocator: HostTensorAllocator, +) -> torch.Tensor: + """ + Allocate tensor and register host memory with cudaHostRegister. + CudaHostRegister only applies when pin_memory=True. + """ + buffer = allocator.allocate(dims, dtype=dtype, device=device) + if pin_memory: + _cuda_host_register(buffer) + return buffer + + +def alloc_with_pin_memory( + dims: tuple, + dtype: torch.dtype, + device: str, + pin_memory: bool, + allocator: None, +) -> torch.Tensor: + """ + Allocate tensor using PyTorch's built-in pin_memory flag. + """ + buffer = torch.empty(dims, dtype=dtype, device=device, pin_memory=pin_memory) + return buffer + + +ALLOC_MEMORY_FUNCS = defaultdict( + lambda: alloc_with_host_register, + { + "npu": alloc_with_pin_memory, + "musa": alloc_with_pin_memory, + }, +) diff --git a/python/sglang/srt/mem_cache/pool_host/hisparse.py b/python/sglang/srt/mem_cache/pool_host/hisparse.py new file mode 100644 index 000000000..fbc9caaa9 --- /dev/null +++ b/python/sglang/srt/mem_cache/pool_host/hisparse.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import logging +from typing import Optional + +import torch + +logger = logging.getLogger(__name__) + + +class HiSparseHostPoolMixin: + def _round_up_to_page_size(self, size: int) -> int: + return (size + self.page_size - 1) // self.page_size * self.page_size + + def alloc_page(self, num_pages: int) -> Optional[torch.Tensor]: + return self.alloc(num_pages * self.page_size) + + def alloc_paged_token_slots( + self, + req_to_host_pool: torch.Tensor, + req_to_host_pool_allocated_len: torch.Tensor, + req_pool_idx: int, + start_pos: int, + num_tokens: int, + ) -> torch.Tensor: + """Allocate request host slots by page and return token-granular slots.""" + device = req_to_host_pool.device + if num_tokens <= 0: + return torch.empty((0,), dtype=torch.int64, device=device) + + allocated_len = int(req_to_host_pool_allocated_len[req_pool_idx]) + end_pos = start_pos + num_tokens + page_end = self._round_up_to_page_size(end_pos) + assert start_pos <= allocated_len + + if page_end > allocated_len: + num_new_pages = (page_end - allocated_len) // self.page_size + host_locs = self.alloc_page(num_new_pages) + if host_locs is None: + logger.error( + "HiSparse: host mem pool alloc failed for %d host pages " + "(req_pool_idx=%d, start_pos=%d, num_tokens=%d)", + num_new_pages, + req_pool_idx, + start_pos, + num_tokens, + ) + raise RuntimeError( + f"HiSparse host mem pool alloc failed for {num_new_pages} pages" + ) + + req_to_host_pool[req_pool_idx, allocated_len:page_end] = host_locs.to( + device=device, non_blocking=True + ) + req_to_host_pool_allocated_len[req_pool_idx] = page_end + + return req_to_host_pool[req_pool_idx, start_pos:end_pos] + + def allocated_host_indices( + self, + req_to_host_pool: torch.Tensor, + req_pool_idx: int, + allocated_len: int, + ) -> torch.Tensor: + allocated_len = int(allocated_len) + host_len = min( + self._round_up_to_page_size(allocated_len), + req_to_host_pool.shape[1], + ) + host_indices = req_to_host_pool[req_pool_idx, :host_len] + return host_indices[host_indices >= 0] diff --git a/python/sglang/srt/mem_cache/storage/aibrix_kvcache/aibrix_kvcache_storage.py b/python/sglang/srt/mem_cache/storage/aibrix_kvcache/aibrix_kvcache_storage.py index bcc827109..95f4d08a4 100644 --- a/python/sglang/srt/mem_cache/storage/aibrix_kvcache/aibrix_kvcache_storage.py +++ b/python/sglang/srt/mem_cache/storage/aibrix_kvcache/aibrix_kvcache_storage.py @@ -18,7 +18,7 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorageConfig, HiCacheStorageExtraInfo, ) -from sglang.srt.mem_cache.memory_pool_host import HostKVCache +from sglang.srt.mem_cache.pool_host import HostKVCache logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/mem_cache/storage/eic/eic_storage.py b/python/sglang/srt/mem_cache/storage/eic/eic_storage.py index f3cc15632..adb1e4675 100644 --- a/python/sglang/srt/mem_cache/storage/eic/eic_storage.py +++ b/python/sglang/srt/mem_cache/storage/eic/eic_storage.py @@ -13,7 +13,7 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorageConfig, HiCacheStorageExtraInfo, ) -from sglang.srt.mem_cache.memory_pool_host import HostKVCache +from sglang.srt.mem_cache.pool_host import HostKVCache logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py b/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py index d9e845e0e..abf07c680 100644 --- a/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py +++ b/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py @@ -22,7 +22,7 @@ from sglang.srt.mem_cache.hicache_storage import ( PoolTransfer, PoolTransferResult, ) -from sglang.srt.mem_cache.memory_pool_host import HostKVCache +from sglang.srt.mem_cache.pool_host import HostKVCache from sglang.srt.mem_cache.storage.hf3fs.hf3fs_client import Hf3fsClient from sglang.srt.observability.metrics_collector import StorageMetrics diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py index 3c1ab73dc..4da1b9f5d 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py @@ -21,11 +21,8 @@ from sglang.srt.mem_cache.hicache_storage import ( PoolTransfer, PoolTransferResult, ) -from sglang.srt.mem_cache.memory_pool_host import ( - HostKVCache, - HostTensorAllocator, - MLATokenToKVPoolHost, -) +from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost +from sglang.srt.mem_cache.pool_host import HostKVCache, HostTensorAllocator from sglang.srt.observability.metrics_collector import StorageMetrics DEFAULT_LOCAL_BUFFER_SIZE = 16 * 1024 * 1024 # 16 MB 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 62d67c7f9..095262477 100644 --- a/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py +++ b/python/sglang/srt/mem_cache/storage/nixl/hicache_nixl.py @@ -13,8 +13,8 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorageConfig, HiCacheStorageExtraInfo, ) -from sglang.srt.mem_cache.memory_pool_host import HostKVCache from sglang.srt.mem_cache.mmap_allocator import alloc_mmap +from sglang.srt.mem_cache.pool_host import HostKVCache from .nixl_registry import NixlRegistry from .nixl_utils import NixlBackendConfig, NixlBackendSelection, NixlFileManager diff --git a/python/sglang/srt/mem_cache/storage/simm/hicache_simm.py b/python/sglang/srt/mem_cache/storage/simm/hicache_simm.py index 8b29fd2d2..66d274ab4 100644 --- a/python/sglang/srt/mem_cache/storage/simm/hicache_simm.py +++ b/python/sglang/srt/mem_cache/storage/simm/hicache_simm.py @@ -16,7 +16,7 @@ from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorageConfig, HiCacheStorageExtraInfo, ) -from sglang.srt.mem_cache.memory_pool_host import HostKVCache +from sglang.srt.mem_cache.pool_host import HostKVCache # Third Party try: diff --git a/test/registered/jit/test_hicache.py b/test/registered/jit/test_hicache.py index 7fa17c4cd..e26b23330 100644 --- a/test/registered/jit/test_hicache.py +++ b/test/registered/jit/test_hicache.py @@ -6,9 +6,11 @@ import torch from sglang.jit_kernel.hicache import can_use_write_back_jit_kernel from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, MLATokenToKVPool from sglang.srt.mem_cache.memory_pool_host import ( - ALLOC_MEMORY_FUNCS, MHATokenToKVPoolHost, MLATokenToKVPoolHost, +) +from sglang.srt.mem_cache.pool_host.common import ( + ALLOC_MEMORY_FUNCS, alloc_with_pin_memory, ) from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu diff --git a/test/registered/unit/managers/test_hisparse_unit.py b/test/registered/unit/managers/test_hisparse_unit.py index 2328cbaef..0cee301a3 100644 --- a/test/registered/unit/managers/test_hisparse_unit.py +++ b/test/registered/unit/managers/test_hisparse_unit.py @@ -85,7 +85,7 @@ class TestHiSparseUnit(unittest.TestCase): torch.distributed.init_process_group(backend="gloo", rank=0, world_size=1) cls.tp_group = torch.distributed.group.WORLD - from sglang.srt.mem_cache.memory_pool_host import ( + from sglang.srt.mem_cache.pool_host.common import ( ALLOC_MEMORY_FUNCS, alloc_with_pin_memory, ) @@ -154,7 +154,7 @@ class TestHiSparseUnit(unittest.TestCase): @classmethod def tearDownClass(cls): - from sglang.srt.mem_cache.memory_pool_host import ALLOC_MEMORY_FUNCS + from sglang.srt.mem_cache.pool_host.common import ALLOC_MEMORY_FUNCS ALLOC_MEMORY_FUNCS["cuda"] = cls._original_alloc if torch.distributed.is_initialized(): diff --git a/test/registered/unit/mem_cache/test_dsa_pool_host_unit.py b/test/registered/unit/mem_cache/test_dsa_pool_host_unit.py index cf553018c..1c6a4825d 100644 --- a/test/registered/unit/mem_cache/test_dsa_pool_host_unit.py +++ b/test/registered/unit/mem_cache/test_dsa_pool_host_unit.py @@ -5,9 +5,11 @@ import torch from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool_host import ( - ALLOC_MEMORY_FUNCS, DSAIndexerPoolHost, MLATokenToKVPoolHost, +) +from sglang.srt.mem_cache.pool_host.common import ( + ALLOC_MEMORY_FUNCS, alloc_with_pin_memory, ) from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu