feat: Support HiCache for MiMo-V2 models (1/N) (#27378)

Co-authored-by: Zhangheng <hzh0425@apache.org>
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Yingchun Lai
2026-06-13 15:40:27 +08:00
committed by GitHub
co-authored by Zhangheng 晟海
parent 60d4bd4c70
commit 806365e778
11 changed files with 667 additions and 43 deletions
@@ -19,8 +19,8 @@ from sglang.srt.mem_cache.memory_pool import (
ReqToTokenPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
get_mha_host_pool_cls,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ceil_align
@@ -57,7 +57,7 @@ class DecodeKVCacheOffloadManager:
)
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
if isinstance(kv_cache, MHATokenToKVPool):
self.decode_host_mem_pool = MHATokenToKVPoolHost(
self.decode_host_mem_pool = get_mha_host_pool_cls(kv_cache)(
kv_cache,
server_args.hicache_ratio,
server_args.hicache_size,
+2 -2
View File
@@ -44,8 +44,8 @@ from sglang.srt.mem_cache.memory_pool import (
MLATokenToKVPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
get_mha_host_pool_cls,
)
from sglang.srt.mem_cache.radix_cache import (
RadixCache,
@@ -78,7 +78,7 @@ class HiRadixCache(RadixCache):
self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache()
if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = MHATokenToKVPoolHost(
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
self.kv_cache,
server_args.hicache_ratio,
server_args.hicache_size,
@@ -19,9 +19,9 @@ from sglang.srt.mem_cache.memory_pool_host import (
HostPoolGroup,
LogicalHostPool,
MambaPoolHost,
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
PoolEntry,
get_mha_host_pool_cls,
)
from sglang.srt.mem_cache.unified_cache_components import ComponentType
@@ -57,7 +57,9 @@ def build_kv_host_pool(
use_mla: bool,
override_kv_cache_dim: Optional[int] = None,
):
kv_host_pool_cls = MLATokenToKVPoolHost if use_mla else MHATokenToKVPoolHost
kv_host_pool_cls = (
MLATokenToKVPoolHost if use_mla else get_mha_host_pool_cls(kv_pool)
)
kwargs = {}
if override_kv_cache_dim is not None:
kwargs["override_kv_cache_dim"] = override_kv_cache_dim
@@ -89,8 +89,8 @@ def maybe_register_hicache_draft(
MLATokenToKVPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
get_mha_host_pool_cls,
)
pool = draft_kv_pool
@@ -107,7 +107,7 @@ def maybe_register_hicache_draft(
layout=server_args.hicache_mem_layout,
)
if isinstance(pool, MHATokenToKVPool):
draft_host_pool = MHATokenToKVPoolHost(pool, **kw)
draft_host_pool = get_mha_host_pool_cls(pool)(pool, **kw)
elif isinstance(pool, MLATokenToKVPool):
draft_host_pool = MLATokenToKVPoolHost(pool, **kw)
else:
+271 -15
View File
@@ -177,8 +177,21 @@ def get_allocator_from_storage(allocator_type):
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,
dims: tuple,
dtype: torch.dtype,
device: str,
pin_memory: bool,
@@ -190,21 +203,12 @@ def alloc_with_host_register(
"""
buffer = allocator.allocate(dims, dtype=dtype, device=device)
if pin_memory:
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."
)
_cuda_host_register(buffer)
return buffer
def alloc_with_pin_memory(
dims,
dims: tuple,
dtype: torch.dtype,
device: str,
pin_memory: bool,
@@ -429,7 +433,6 @@ class MHATokenToKVPoolHost(HostKVCache):
self.head_num = self.device_pool.head_num
self.head_dim = self.device_pool.head_dim
self.layer_num = self.device_pool.layer_num
return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize * 2
def get_ksize_per_token(self):
@@ -810,7 +813,7 @@ class MHATokenToKVPoolHost(HostKVCache):
return ptr_list, element_size_list
def get_page_buffer_meta(self, indices):
""" "
"""
meta data for zero copy
"""
assert len(indices) % self.page_size == 0
@@ -896,6 +899,259 @@ class MHATokenToKVPoolHost(HostKVCache):
return base_aligned and stride % page_size_bytes == 0
class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
"""Host KV pool for MHA models whose K and V have different head dims
(``head_dim != v_head_dim``), e.g. MiMo-V2.
K and V are stored in two independent host buffers (``self.k_buffer`` and
``self.v_buffer``) instead of a single ``(2, ...)`` tensor, so each side
keeps its native stride. The kernel transfer path dispatches K and V as
independent single-buffer copies so each side uses its own ``item_size``.
Direct transfer and the flat-page L3 storage interface assume a single
shared ``item_size`` in paths that are not safe for asymmetric K/V, so they
raise instead of silently corrupting V copies.
"""
def get_size_per_token(self):
self.head_num = self.device_pool.head_num
self.head_dim = self.device_pool.head_dim
self.layer_num = self.device_pool.layer_num
self.v_head_dim = self.device_pool.v_head_dim
return (
(self.head_dim + self.v_head_dim)
* self.head_num
* self.layer_num
* self.dtype.itemsize
)
def get_ksize_per_token(self):
return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize
def init_kv_buffer(self):
if self.layout == "page_first":
k_dims = (self.size, self.layer_num, self.head_num, self.head_dim)
v_dims = (self.size, self.layer_num, self.head_num, self.v_head_dim)
else:
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim: "
f"{self.layout}; expected 'page_first'."
)
# token_stride_size / layout_dim are intentionally NOT set: K and V
# have different strides, so any caller that reaches for a single
# shared stride is a bug. Such callers will fail loudly with
# AttributeError rather than silently use the K stride for V copies.
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
k_buffer = alloc_func(
k_dims,
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
)
v_buffer = alloc_func(
v_dims,
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
)
return (k_buffer, v_buffer)
def _k_token_stride_size(self) -> int:
return self.head_num * self.head_dim * self.dtype.itemsize
def _v_token_stride_size(self) -> int:
return self.head_num * self.v_head_dim * self.dtype.itemsize
def _k_layout_dim(self) -> int:
return self._k_token_stride_size() * self.layer_num
def _v_layout_dim(self) -> int:
return self._v_token_stride_size() * self.layer_num
def _flat_page_unsupported(self) -> NotImplementedError:
return NotImplementedError(
"Models with head_dim != v_head_dim do not support the flat-page "
"interface used by HiCache L3 storage backends {hf3fs, eic, nixl}. "
"Use a backend that does not use this interface (e.g. mooncake, simm)."
)
def load_to_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
):
if io_backend == "kernel":
if self.layout != "page_first":
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim "
f"and io_backend='kernel': {self.layout}; expected 'page_first'."
)
transfer_kv_per_layer_mla_pf_lf(
src=self.k_buffer,
dst=device_pool.k_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
item_size=self._k_token_stride_size(),
src_layout_dim=self._k_layout_dim(),
)
transfer_kv_per_layer_mla_pf_lf(
src=self.v_buffer,
dst=device_pool.v_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
item_size=self._v_token_stride_size(),
src_layout_dim=self._v_layout_dim(),
)
else:
raise ValueError(
f"Unsupported IO backend for models with head_dim != v_head_dim: "
f"{io_backend}; expected 'kernel'."
)
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
if io_backend == "kernel":
if self.layout != "page_first":
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim "
f"and io_backend='kernel': {self.layout}; expected 'page_first'."
)
transfer_kv_all_layer_mla_lf_pf(
src_layers=device_pool.k_data_ptrs,
dst=self.k_buffer,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self._k_token_stride_size(),
dst_layout_dim=self._k_layout_dim(),
num_layers=self.layer_num,
)
transfer_kv_all_layer_mla_lf_pf(
src_layers=device_pool.v_data_ptrs,
dst=self.v_buffer,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self._v_token_stride_size(),
dst_layout_dim=self._v_layout_dim(),
num_layers=self.layer_num,
)
else:
raise ValueError(
f"Unsupported IO backend for models with head_dim != v_head_dim: "
f"{io_backend}; expected 'kernel'."
)
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
raise self._flat_page_unsupported()
def get_dummy_flat_data_page(self) -> torch.Tensor:
raise self._flat_page_unsupported()
def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None:
raise self._flat_page_unsupported()
def get_split_heads_page_buffer_meta(
self, indices: torch.Tensor, split_factor: int
):
raise NotImplementedError(
"get_split_heads_page_buffer_meta requires layout='page_head', "
"which is not supported for models with head_dim != v_head_dim."
)
def get_page_buffer_meta(self, indices):
assert len(indices) % self.page_size == 0
if self.layout != "page_first":
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim: "
f"{self.layout}"
)
indices = indices.tolist()
k_base_ptr = self.k_buffer.data_ptr()
v_base_ptr = self.v_buffer.data_ptr()
k_element_size = (
self.layer_num
* self.dtype.itemsize
* self.page_size
* self.head_num
* self.head_dim
)
v_element_size = (
self.layer_num
* self.dtype.itemsize
* self.page_size
* self.head_num
* self.v_head_dim
)
ptr_list = []
element_size_list = []
for index in range(0, len(indices), self.page_size):
k_ptr = (
k_base_ptr
+ indices[index]
* self.layer_num
* self.head_num
* self.head_dim
* self.dtype.itemsize
)
v_ptr = (
v_base_ptr
+ indices[index]
* self.layer_num
* self.head_num
* self.v_head_dim
* self.dtype.itemsize
)
ptr_list.extend([k_ptr, v_ptr])
element_size_list.extend([k_element_size, v_element_size])
return ptr_list, element_size_list
def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool:
if self.layout != "page_first":
return False
k_stride = (
self.page_size
* self.layer_num
* self.head_num
* self.head_dim
* self.dtype.itemsize
)
v_stride = (
self.page_size
* self.layer_num
* self.head_num
* self.v_head_dim
* self.dtype.itemsize
)
base_aligned = (
self.k_buffer.data_ptr() % page_size_bytes == 0
and self.v_buffer.data_ptr() % page_size_bytes == 0
)
return (
base_aligned
and k_stride % page_size_bytes == 0
and v_stride % page_size_bytes == 0
)
def get_mha_host_pool_cls(device_pool: MHATokenToKVPool) -> type:
"""Pick the right MHA host-pool class based on the device pool's K/V dims.
Returns ``AsymmetricMHATokenToKVPoolHost`` when ``head_dim != v_head_dim``
(e.g. MiMo-V2), else the default ``MHATokenToKVPoolHost``.
"""
if device_pool.head_dim != device_pool.v_head_dim:
return AsymmetricMHATokenToKVPoolHost
return MHATokenToKVPoolHost
class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
device_pool: MLATokenToKVPool
@@ -1256,7 +1512,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
raise ValueError(f"Unsupported layout: {self.layout}")
def get_page_buffer_meta(self, indices):
""" "
"""
meta data for zero copy
"""
assert len(indices) % self.page_size == 0
+23 -8
View File
@@ -2418,14 +2418,29 @@ class ServerArgs:
)
if self.enable_hierarchical_cache:
self.swa_full_tokens_ratio = 1.0
logger.warning(
"Reset swa_full_tokens_ratio to 1.0 for MiMoV2 model with hierarchical cache"
)
self.disable_hybrid_swa_memory = True
logger.warning(
"Disable hybrid SWA memory for MiMoV2 model with hierarchical cache"
)
if not envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get():
raise ValueError(
"Hierarchical cache for MiMoV2 requires the unified "
"radix tree. Set SGLANG_ENABLE_UNIFIED_RADIX_TREE=1 "
"to enable --enable-hierarchical-cache for this model."
)
# MiMoV2 has head_dim != v_head_dim, so the host KV pool uses
# asymmetric K/V allocation. Only the kernel/page_first transfer
# path has a safe split K/V implementation.
if self.hicache_io_backend != "kernel":
logger.warning(
f"Force hicache_io_backend to 'kernel' for MiMoV2 model "
f"(was {self.hicache_io_backend!r})."
)
self.hicache_io_backend = "kernel"
if self.hicache_mem_layout != "page_first":
logger.warning(
f"Force hicache_mem_layout to 'page_first' for "
f"MiMoV2 model (was {self.hicache_mem_layout!r}); "
f"asymmetric K/V HiCache requires kernel/page_first."
)
self.hicache_mem_layout = "page_first"
elif (
"Step3p5ForCausalLM" in model_arch
or "Step3p7ForConditionalGeneration" in model_arch