Files
sglang/python/sglang/srt/mem_cache/hisparse_memory_pool.py
T

258 lines
8.9 KiB
Python

# mapping on device memory, host memory and memory allocator
import logging
from typing import Optional
import torch
from sglang.kernels.ops.kvcache.hisparse_slot_mapping import (
translate_padded_hisparse_locations,
)
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, MHATokenToKVPool
from sglang.srt.utils import is_cuda, is_hip, is_xpu
logger = logging.getLogger(__name__)
# sgl_kernel.kvcacheio is only available in CUDA/ROCm/XPU sgl-kernel builds (not MPS/NPU/CPU).
_is_cuda = is_cuda()
_is_hip = is_hip()
_is_xpu = is_xpu()
if _is_cuda or _is_hip or _is_xpu:
from sgl_kernel.kvcacheio import transfer_kv_all_layer_mla
else:
def transfer_kv_all_layer_mla(*args, **kwargs):
raise RuntimeError(
"HiSparse device KV transfer requires sgl_kernel.kvcacheio "
"(CUDA/ROCm/XPU). It is not available on this backend."
)
class HiSparseDSATokenToKVPool(DSATokenToKVPool):
def __init__(
self,
size: int,
page_size: int,
kv_lora_rank: int,
dtype: torch.dtype,
qk_rope_head_dim: int,
layer_num: int,
device: str,
index_head_dim: int,
enable_memory_saver: bool,
kv_cache_dim: int,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
index_kpool: int = 1,
index_kpool_compress: bool = False,
tail_extra_slots: int = 0,
max_running_requests: Optional[int] = None,
skip_topk_layers: Optional[list[bool]] = None,
host_to_device_ratio: int = 2,
):
super().__init__(
size=size,
page_size=page_size,
kv_lora_rank=kv_lora_rank,
dtype=dtype,
qk_rope_head_dim=qk_rope_head_dim,
layer_num=layer_num,
device=device,
index_head_dim=index_head_dim,
enable_memory_saver=enable_memory_saver,
kv_cache_dim=kv_cache_dim,
start_layer=start_layer,
end_layer=end_layer,
index_buf_size=size * host_to_device_ratio,
index_kpool=index_kpool,
index_kpool_compress=index_kpool_compress,
tail_extra_slots=tail_extra_slots,
max_running_requests=max_running_requests,
skip_topk_layers=skip_topk_layers,
)
self.bytes_per_token = self.kv_cache_dim * self.dtype.itemsize
def register_mapping(self, full_to_hisparse_device_index_mapping: torch.Tensor):
self.full_to_hisparse_device_index_mapping = (
full_to_hisparse_device_index_mapping
)
def translate_loc_to_hisparse_device(
self, compressed_indices: torch.Tensor
) -> torch.Tensor:
"""Map logical locations to physical slots with the same shape.
CUDA and ROCm use a fused kernel for 1D GPU slot lists, preserving
negative padding. Page tables and CPU inputs keep the direct gather.
"""
if compressed_indices.is_cuda and compressed_indices.ndim == 1:
return translate_padded_hisparse_locations(
self.full_to_hisparse_device_index_mapping, compressed_indices
)
return self.full_to_hisparse_device_index_mapping[compressed_indices]
def _translate_loc_to_hisparse_device(self, compressed_indices: torch.Tensor):
return self.full_to_hisparse_device_index_mapping[compressed_indices]
def translate_loc_from_full_to_hisparse_device(self, full_indices: torch.Tensor):
return self._translate_loc_to_hisparse_device(full_indices)
def translate_loc_from_full_to_compressed(self, full_indices: torch.Tensor):
return full_indices
def set_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
):
loc = self.translate_loc_to_hisparse_device(loc)
super().set_kv_buffer(layer, loc, cache_k, cache_v)
def set_mla_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_k_nope: torch.Tensor,
cache_k_rope: torch.Tensor,
):
loc = self.translate_loc_to_hisparse_device(loc)
super().set_mla_kv_buffer(layer, loc, cache_k_nope, cache_k_rope)
def get_mla_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
dst_dtype: Optional[torch.dtype] = None,
):
loc = self.translate_loc_to_hisparse_device(loc)
return super().get_mla_kv_buffer(layer, loc, dst_dtype)
def transfer_values_on_device(self, dst_indices, src_indices):
transfer_kv_all_layer_mla(
src_layers=self.data_ptrs,
dst_layers=self.data_ptrs,
src_indices=src_indices,
dst_indices=dst_indices,
item_size=self.bytes_per_token,
num_layers=self.layer_num,
)
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
raise NotImplementedError("HiSparseDevicePool does not support get_cpu_copy")
def load_cpu_copy(
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
):
raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy")
class HiSparseMHAMainPool(MHATokenToKVPool):
"""MHA KV pool with HiSparse logical-to-device mapping.
Used by MiniMax M3 HiSparse. The index pools (index_kv_pool, index_k_pool)
stay fully resident on the device and do not use this mapping.
"""
def __init__(
self,
size: int,
page_size: int,
dtype: torch.dtype,
head_num: int,
head_dim: int,
layer_num: int,
device: str,
enable_memory_saver: bool,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
):
super().__init__(
size=size,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=layer_num,
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=end_layer,
)
self.full_to_hisparse_device_index_mapping: Optional[torch.Tensor] = None
self.bytes_per_token_k = head_num * head_dim * self.store_dtype.itemsize
self.bytes_per_token_v = head_num * self.v_head_dim * self.store_dtype.itemsize
def register_mapping(
self, full_to_hisparse_device_index_mapping: torch.Tensor
) -> None:
self.full_to_hisparse_device_index_mapping = (
full_to_hisparse_device_index_mapping
)
def translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[indices]
def _translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[indices]
def translate_loc_from_full_to_hisparse_device(
self, full_indices: torch.Tensor
) -> torch.Tensor:
assert self.full_to_hisparse_device_index_mapping is not None
return self.full_to_hisparse_device_index_mapping[full_indices]
def translate_loc_from_full_to_compressed(
self, full_indices: torch.Tensor
) -> torch.Tensor:
return full_indices
def set_kv_buffer(
self,
layer: RadixAttention,
loc,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
*args,
**kwargs,
):
from sglang.srt.mem_cache.memory_pool import unwrap_write_loc
raw_loc, _, _ = unwrap_write_loc(loc)
translated = self.translate_loc_to_hisparse_device(raw_loc)
super().set_kv_buffer(layer, translated, cache_k, cache_v, *args, **kwargs)
def transfer_values_on_device(
self,
dst_indices: torch.Tensor,
src_indices: torch.Tensor,
) -> None:
transfer_kv_all_layer_mla(
src_layers=self.k_data_ptrs,
dst_layers=self.k_data_ptrs,
src_indices=src_indices,
dst_indices=dst_indices,
item_size=self.bytes_per_token_k,
num_layers=self.layer_num,
)
transfer_kv_all_layer_mla(
src_layers=self.v_data_ptrs,
dst_layers=self.v_data_ptrs,
src_indices=src_indices,
dst_indices=dst_indices,
item_size=self.bytes_per_token_v,
num_layers=self.layer_num,
)
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
raise NotImplementedError("HiSparseMHAMainPool does not support get_cpu_copy")
def load_cpu_copy(
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
):
raise NotImplementedError("HiSparseMHAMainPool does not support load_cpu_copy")