[mem_cache][5/N] refactor: extract host KV cache base layer into pool_host package (#27273)
This commit is contained in:
@@ -31,7 +31,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
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 (
|
from sglang.srt.distributed import (
|
||||||
get_pipeline_model_parallel_rank,
|
get_pipeline_model_parallel_rank,
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ import torch
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import abc
|
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from collections import defaultdict
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from functools import wraps
|
|
||||||
from typing import TYPE_CHECKING, Any, Callable, Optional
|
from typing import TYPE_CHECKING, Any, Callable, Optional
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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.jit_kernel.hisparse import transfer_cache_dsv4_mla
|
||||||
from sglang.srt.mem_cache.memory_pool import (
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
DSATokenToKVPool,
|
DSATokenToKVPool,
|
||||||
KVCache,
|
|
||||||
MambaPool,
|
MambaPool,
|
||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
MLATokenToKVPool,
|
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
|
from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -74,321 +69,21 @@ if _is_npu:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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
|
_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):
|
class MHATokenToKVPoolHost(HostKVCache):
|
||||||
device_pool: MHATokenToKVPool
|
device_pool: MHATokenToKVPool
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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)
|
||||||
@@ -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,
|
||||||
|
},
|
||||||
|
)
|
||||||
@@ -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]
|
||||||
@@ -18,7 +18,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
HiCacheStorageConfig,
|
HiCacheStorageConfig,
|
||||||
HiCacheStorageExtraInfo,
|
HiCacheStorageExtraInfo,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
|
from sglang.srt.mem_cache.pool_host import HostKVCache
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
HiCacheStorageConfig,
|
HiCacheStorageConfig,
|
||||||
HiCacheStorageExtraInfo,
|
HiCacheStorageExtraInfo,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
|
from sglang.srt.mem_cache.pool_host import HostKVCache
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
PoolTransfer,
|
PoolTransfer,
|
||||||
PoolTransferResult,
|
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.mem_cache.storage.hf3fs.hf3fs_client import Hf3fsClient
|
||||||
from sglang.srt.observability.metrics_collector import StorageMetrics
|
from sglang.srt.observability.metrics_collector import StorageMetrics
|
||||||
|
|
||||||
|
|||||||
@@ -21,11 +21,8 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
PoolTransfer,
|
PoolTransfer,
|
||||||
PoolTransferResult,
|
PoolTransferResult,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||||
HostKVCache,
|
from sglang.srt.mem_cache.pool_host import HostKVCache, HostTensorAllocator
|
||||||
HostTensorAllocator,
|
|
||||||
MLATokenToKVPoolHost,
|
|
||||||
)
|
|
||||||
from sglang.srt.observability.metrics_collector import StorageMetrics
|
from sglang.srt.observability.metrics_collector import StorageMetrics
|
||||||
|
|
||||||
DEFAULT_LOCAL_BUFFER_SIZE = 16 * 1024 * 1024 # 16 MB
|
DEFAULT_LOCAL_BUFFER_SIZE = 16 * 1024 * 1024 # 16 MB
|
||||||
|
|||||||
@@ -13,8 +13,8 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
HiCacheStorageConfig,
|
HiCacheStorageConfig,
|
||||||
HiCacheStorageExtraInfo,
|
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.mmap_allocator import alloc_mmap
|
||||||
|
from sglang.srt.mem_cache.pool_host import HostKVCache
|
||||||
|
|
||||||
from .nixl_registry import NixlRegistry
|
from .nixl_registry import NixlRegistry
|
||||||
from .nixl_utils import NixlBackendConfig, NixlBackendSelection, NixlFileManager
|
from .nixl_utils import NixlBackendConfig, NixlBackendSelection, NixlFileManager
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
HiCacheStorageConfig,
|
HiCacheStorageConfig,
|
||||||
HiCacheStorageExtraInfo,
|
HiCacheStorageExtraInfo,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
|
from sglang.srt.mem_cache.pool_host import HostKVCache
|
||||||
|
|
||||||
# Third Party
|
# Third Party
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -6,9 +6,11 @@ import torch
|
|||||||
from sglang.jit_kernel.hicache import can_use_write_back_jit_kernel
|
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 import MHATokenToKVPool, MLATokenToKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
ALLOC_MEMORY_FUNCS,
|
|
||||||
MHATokenToKVPoolHost,
|
MHATokenToKVPoolHost,
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.common import (
|
||||||
|
ALLOC_MEMORY_FUNCS,
|
||||||
alloc_with_pin_memory,
|
alloc_with_pin_memory,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ class TestHiSparseUnit(unittest.TestCase):
|
|||||||
torch.distributed.init_process_group(backend="gloo", rank=0, world_size=1)
|
torch.distributed.init_process_group(backend="gloo", rank=0, world_size=1)
|
||||||
cls.tp_group = torch.distributed.group.WORLD
|
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_MEMORY_FUNCS,
|
||||||
alloc_with_pin_memory,
|
alloc_with_pin_memory,
|
||||||
)
|
)
|
||||||
@@ -154,7 +154,7 @@ class TestHiSparseUnit(unittest.TestCase):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def tearDownClass(cls):
|
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
|
ALLOC_MEMORY_FUNCS["cuda"] = cls._original_alloc
|
||||||
if torch.distributed.is_initialized():
|
if torch.distributed.is_initialized():
|
||||||
|
|||||||
@@ -5,9 +5,11 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
ALLOC_MEMORY_FUNCS,
|
|
||||||
DSAIndexerPoolHost,
|
DSAIndexerPoolHost,
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.common import (
|
||||||
|
ALLOC_MEMORY_FUNCS,
|
||||||
alloc_with_pin_memory,
|
alloc_with_pin_memory,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||||
|
|||||||
Reference in New Issue
Block a user