258 lines
8.9 KiB
Python
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")
|