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:
co-authored by
Zhangheng
晟海
parent
60d4bd4c70
commit
806365e778
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool_host import AsymmetricMHATokenToKVPoolHost
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large")
|
||||
|
||||
# These tests use AsymmetricMHATokenToKVPoolHost methods and let that class call
|
||||
# the real sgl-kernel transfer ops. The asymmetric host pool is kernel-only;
|
||||
# direct/page_first_direct is intentionally rejected in the CPU dispatch tests.
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(), reason="asymmetric host-pool tests require CUDA."
|
||||
)
|
||||
|
||||
DEVICE = "cuda"
|
||||
PAGE_SIZE = 16
|
||||
NUM_LAYERS = 3
|
||||
TOTAL_ITEMS = PAGE_SIZE * 8
|
||||
HEAD_NUM = 4
|
||||
K_HEAD_DIM = 192
|
||||
V_HEAD_DIM = 128
|
||||
DTYPES = [torch.float16, torch.bfloat16]
|
||||
|
||||
|
||||
def token_indices_for_pages(pages, page_size=PAGE_SIZE, device=None):
|
||||
indices = torch.cat(
|
||||
[
|
||||
torch.arange(
|
||||
int(page) * page_size,
|
||||
(int(page) + 1) * page_size,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
for page in pages.tolist()
|
||||
]
|
||||
)
|
||||
return indices if device is None else indices.to(device)
|
||||
|
||||
|
||||
def fill_with_offset(tensor, offset):
|
||||
data = torch.arange(tensor.numel(), device=tensor.device, dtype=tensor.dtype)
|
||||
tensor.copy_((data + offset).view_as(tensor))
|
||||
|
||||
|
||||
def make_host_pool(dtype):
|
||||
host = AsymmetricMHATokenToKVPoolHost.__new__(AsymmetricMHATokenToKVPoolHost)
|
||||
host.layout = "page_first"
|
||||
host.page_size = PAGE_SIZE
|
||||
host.layer_num = NUM_LAYERS
|
||||
host.head_num = HEAD_NUM
|
||||
host.head_dim = K_HEAD_DIM
|
||||
host.v_head_dim = V_HEAD_DIM
|
||||
host.dtype = dtype
|
||||
host.kv_buffer = (
|
||||
torch.zeros(
|
||||
TOTAL_ITEMS, NUM_LAYERS, HEAD_NUM, K_HEAD_DIM, dtype=dtype
|
||||
).pin_memory(),
|
||||
torch.zeros(
|
||||
TOTAL_ITEMS, NUM_LAYERS, HEAD_NUM, V_HEAD_DIM, dtype=dtype
|
||||
).pin_memory(),
|
||||
)
|
||||
return host
|
||||
|
||||
|
||||
def make_device_pool(dtype):
|
||||
k_buffer = [
|
||||
torch.empty(TOTAL_ITEMS, HEAD_NUM, K_HEAD_DIM, dtype=dtype, device=DEVICE)
|
||||
for _ in range(NUM_LAYERS)
|
||||
]
|
||||
v_buffer = [
|
||||
torch.empty(TOTAL_ITEMS, HEAD_NUM, V_HEAD_DIM, dtype=dtype, device=DEVICE)
|
||||
for _ in range(NUM_LAYERS)
|
||||
]
|
||||
for layer_id in range(NUM_LAYERS):
|
||||
fill_with_offset(k_buffer[layer_id], layer_id * 1000)
|
||||
fill_with_offset(v_buffer[layer_id], layer_id * 1000 + 100)
|
||||
|
||||
return SimpleNamespace(
|
||||
k_buffer=k_buffer,
|
||||
v_buffer=v_buffer,
|
||||
k_data_ptrs=torch.tensor(
|
||||
[x.data_ptr() for x in k_buffer], dtype=torch.uint64, device=DEVICE
|
||||
),
|
||||
v_data_ptrs=torch.tensor(
|
||||
[x.data_ptr() for x in v_buffer], dtype=torch.uint64, device=DEVICE
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def assert_backup_matches_device(host, device_pool, host_indices_host, device_indices):
|
||||
for layer_id in range(NUM_LAYERS):
|
||||
torch.testing.assert_close(
|
||||
host.k_buffer[host_indices_host, layer_id],
|
||||
device_pool.k_buffer[layer_id][device_indices].cpu(),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
host.v_buffer[host_indices_host, layer_id],
|
||||
device_pool.v_buffer[layer_id][device_indices].cpu(),
|
||||
)
|
||||
|
||||
|
||||
def assert_load_matches_host(host, device_pool, host_indices_host, load_indices):
|
||||
for layer_id in range(NUM_LAYERS):
|
||||
torch.testing.assert_close(
|
||||
device_pool.k_buffer[layer_id][load_indices],
|
||||
host.k_buffer[host_indices_host, layer_id].to(DEVICE),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
device_pool.v_buffer[layer_id][load_indices],
|
||||
host.v_buffer[host_indices_host, layer_id].to(DEVICE),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", DTYPES)
|
||||
def test_asymmetric_mha_kernel_page_first_roundtrip(dtype):
|
||||
# Covers D2H backup + H2D load through AsymmetricMHATokenToKVPoolHost using
|
||||
# MiMoV2's real K/V head dims and the real MLA single-buffer kernels.
|
||||
host = make_host_pool(dtype)
|
||||
device_pool = make_device_pool(dtype)
|
||||
|
||||
device_pages = torch.tensor([1, 2, 3], dtype=torch.int64)
|
||||
host_pages = torch.tensor([0, 1, 2], dtype=torch.int64)
|
||||
load_pages = torch.tensor([4, 5, 6], dtype=torch.int64)
|
||||
device_indices_host = token_indices_for_pages(device_pages)
|
||||
host_indices_host = token_indices_for_pages(host_pages)
|
||||
load_indices_host = token_indices_for_pages(load_pages)
|
||||
device_indices = device_indices_host.to(DEVICE)
|
||||
host_indices = host_indices_host.to(DEVICE)
|
||||
load_indices = load_indices_host.to(DEVICE)
|
||||
|
||||
host.backup_from_device_all_layer(
|
||||
device_pool, host_indices, device_indices, io_backend="kernel"
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
assert_backup_matches_device(
|
||||
host, device_pool, host_indices_host, device_indices_host
|
||||
)
|
||||
|
||||
for layer_id in range(NUM_LAYERS):
|
||||
device_pool.k_buffer[layer_id].zero_()
|
||||
device_pool.v_buffer[layer_id].zero_()
|
||||
host.load_to_device_per_layer(
|
||||
device_pool, host_indices, load_indices, layer_id, io_backend="kernel"
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
assert_load_matches_host(host, device_pool, host_indices_host, load_indices_host)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||
@@ -1,5 +1,6 @@
|
||||
import unittest
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||
from sglang.test.server_fixtures.mmmu_fixture import MMMUServerBase
|
||||
@@ -20,6 +21,13 @@ MIMO_V2_OTHER_ARGS = [
|
||||
"fa3",
|
||||
"--reasoning-parser",
|
||||
"mimo",
|
||||
"--enable-hierarchical-cache",
|
||||
"--hicache-ratio",
|
||||
"1.5",
|
||||
"--hicache-mem-layout",
|
||||
"page_first",
|
||||
"--hicache-io-backend",
|
||||
"kernel",
|
||||
]
|
||||
MIMO_V2_MTP_OTHER_ARGS = MIMO_V2_OTHER_ARGS + [
|
||||
"--speculative-algorithm",
|
||||
@@ -42,6 +50,11 @@ class TestMiMoV2(GSM8KMixin, MMMUServerBase):
|
||||
server_api_key = None
|
||||
other_args = MIMO_V2_MTP_OTHER_ARGS
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
with envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.override(True):
|
||||
super().setUpClass()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import unittest
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||
from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin
|
||||
@@ -42,11 +43,23 @@ class TestMiMoV2Flash(GSM8KMixin, SpecDecodingMixin, DefaultServerBase):
|
||||
"--enable-multi-layer-eagle",
|
||||
"--model-loader-extra-config",
|
||||
'{"enable_multithread_load": true,"num_threads": 64}',
|
||||
"--enable-hierarchical-cache",
|
||||
"--hicache-ratio",
|
||||
"1.5",
|
||||
"--hicache-mem-layout",
|
||||
"page_first",
|
||||
"--hicache-io-backend",
|
||||
"kernel",
|
||||
]
|
||||
|
||||
bs_1_speed_thres = 170
|
||||
accept_length_thres = 3.2
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
with envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.override(True):
|
||||
super().setUpClass()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
"""Unit tests for asymmetric MHA host KV pool transfer dispatch."""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
AsymmetricMHATokenToKVPoolHost,
|
||||
MHATokenToKVPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _make_host(layout: str) -> AsymmetricMHATokenToKVPoolHost:
|
||||
host = AsymmetricMHATokenToKVPoolHost.__new__(AsymmetricMHATokenToKVPoolHost)
|
||||
host.layout = layout
|
||||
host.page_size = 2
|
||||
host.layer_num = 3
|
||||
host.head_num = 2
|
||||
host.head_dim = 4
|
||||
host.v_head_dim = 6
|
||||
host.dtype = torch.float16
|
||||
|
||||
if layout == "page_first":
|
||||
k_dims = (8, host.layer_num, host.head_num, host.head_dim)
|
||||
v_dims = (8, host.layer_num, host.head_num, host.v_head_dim)
|
||||
else:
|
||||
raise ValueError(f"Unsupported test layout: {layout}")
|
||||
|
||||
host.kv_buffer = (torch.empty(k_dims), torch.empty(v_dims))
|
||||
return host
|
||||
|
||||
|
||||
def _make_device_pool(host: AsymmetricMHATokenToKVPoolHost) -> SimpleNamespace:
|
||||
size = 8
|
||||
k_buffer = [
|
||||
torch.empty(size, host.head_num, host.head_dim) for _ in range(host.layer_num)
|
||||
]
|
||||
v_buffer = [
|
||||
torch.empty(size, host.head_num, host.v_head_dim) for _ in range(host.layer_num)
|
||||
]
|
||||
return SimpleNamespace(
|
||||
k_buffer=k_buffer,
|
||||
v_buffer=v_buffer,
|
||||
k_data_ptrs=torch.tensor([x.data_ptr() for x in k_buffer], dtype=torch.uint64),
|
||||
v_data_ptrs=torch.tensor([x.data_ptr() for x in v_buffer], dtype=torch.uint64),
|
||||
)
|
||||
|
||||
|
||||
class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
||||
def test_factory_selects_asymmetric_pool_for_mismatched_kv_dims(self):
|
||||
symmetric_pool = SimpleNamespace(head_dim=4, v_head_dim=4)
|
||||
asymmetric_pool = SimpleNamespace(head_dim=4, v_head_dim=6)
|
||||
|
||||
self.assertIs(get_mha_host_pool_cls(symmetric_pool), MHATokenToKVPoolHost)
|
||||
self.assertIs(
|
||||
get_mha_host_pool_cls(asymmetric_pool), AsymmetricMHATokenToKVPoolHost
|
||||
)
|
||||
|
||||
def test_kernel_load_splits_k_and_v_with_separate_strides(self):
|
||||
# Dispatch-only test: the CUDA kernel is mocked; this verifies that K and
|
||||
# V are sent as separate single-buffer calls with their own byte strides.
|
||||
host = _make_host("page_first")
|
||||
device_pool = _make_device_pool(host)
|
||||
host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_per_layer_mla_pf_lf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.load_to_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=1,
|
||||
io_backend="kernel",
|
||||
)
|
||||
|
||||
self.assertEqual(transfer.call_count, 2)
|
||||
k_call, v_call = transfer.call_args_list
|
||||
self.assertIs(k_call.kwargs["src"], host.k_buffer)
|
||||
self.assertIs(k_call.kwargs["dst"], device_pool.k_buffer[1])
|
||||
self.assertEqual(k_call.kwargs["item_size"], 16)
|
||||
self.assertEqual(k_call.kwargs["src_layout_dim"], 48)
|
||||
self.assertIs(v_call.kwargs["src"], host.v_buffer)
|
||||
self.assertIs(v_call.kwargs["dst"], device_pool.v_buffer[1])
|
||||
self.assertEqual(v_call.kwargs["item_size"], 24)
|
||||
self.assertEqual(v_call.kwargs["src_layout_dim"], 72)
|
||||
|
||||
def test_kernel_backup_splits_k_and_v_with_separate_strides(self):
|
||||
# Dispatch-only test: D2H backup must pass separate K/V layer pointer
|
||||
# tables so the single-buffer MLA kernel gets the correct stride per side.
|
||||
host = _make_host("page_first")
|
||||
device_pool = _make_device_pool(host)
|
||||
host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_all_layer_mla_lf_pf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.backup_from_device_all_layer(
|
||||
device_pool, host_indices, device_indices, io_backend="kernel"
|
||||
)
|
||||
|
||||
self.assertEqual(transfer.call_count, 2)
|
||||
k_call, v_call = transfer.call_args_list
|
||||
self.assertIs(k_call.kwargs["src_layers"], device_pool.k_data_ptrs)
|
||||
self.assertIs(k_call.kwargs["dst"], host.k_buffer)
|
||||
self.assertEqual(k_call.kwargs["item_size"], 16)
|
||||
self.assertEqual(k_call.kwargs["dst_layout_dim"], 48)
|
||||
self.assertIs(v_call.kwargs["src_layers"], device_pool.v_data_ptrs)
|
||||
self.assertIs(v_call.kwargs["dst"], host.v_buffer)
|
||||
self.assertEqual(v_call.kwargs["item_size"], 24)
|
||||
self.assertEqual(v_call.kwargs["dst_layout_dim"], 72)
|
||||
|
||||
def test_direct_load_is_rejected(self):
|
||||
# Direct single-buffer D2H is not reliable for asymmetric K/V in the
|
||||
# current sgl-kernel fast path, so the asymmetric host pool is kernel-only.
|
||||
host = _make_host("page_first")
|
||||
device_pool = _make_device_pool(host)
|
||||
host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "expected 'kernel'"):
|
||||
host.load_to_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=2,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
def test_direct_backup_is_rejected(self):
|
||||
# Same restriction for D2H backup: asymmetric MHA uses the kernel path
|
||||
# until the direct kernel has an explicit safe asymmetric mode.
|
||||
host = _make_host("page_first")
|
||||
device_pool = _make_device_pool(host)
|
||||
host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "expected 'kernel'"):
|
||||
host.backup_from_device_all_layer(
|
||||
device_pool, host_indices, device_indices, io_backend="direct"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -450,16 +450,24 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
def _init_hicache(self, tree, *, write_policy: str = "write_through"):
|
||||
import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler
|
||||
|
||||
orig_kv_host_pool = assembler.MHATokenToKVPoolHost
|
||||
# Wrap the host-pool factory (not MHATokenToKVPoolHost directly)
|
||||
# because the assembler picks between MHATokenToKVPoolHost and
|
||||
# AsymmetricMHATokenToKVPoolHost via get_mha_host_pool_cls(device_pool).
|
||||
orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return orig_kv_host_pool(*args, **kwargs)
|
||||
def get_mha_host_pool_cls_wrapper(device_pool):
|
||||
host_pool_cls = orig_get_mha_host_pool_cls(device_pool)
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return host_pool_cls(*args, **kwargs)
|
||||
|
||||
return kv_host_pool_wrapper
|
||||
|
||||
patcher = mock.patch.object(
|
||||
assembler,
|
||||
"MHATokenToKVPoolHost",
|
||||
side_effect=kv_host_pool_wrapper,
|
||||
"get_mha_host_pool_cls",
|
||||
side_effect=get_mha_host_pool_cls_wrapper,
|
||||
)
|
||||
patcher.start()
|
||||
self.addCleanup(patcher.stop)
|
||||
@@ -2370,12 +2378,20 @@ class UnifiedRadixCacheSuite:
|
||||
):
|
||||
import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler
|
||||
|
||||
orig_kv_host_pool = assembler.MHATokenToKVPoolHost
|
||||
# See _init_hicache: wrap the factory rather than MHATokenToKVPoolHost
|
||||
# directly so the pin_memory=False override applies to both
|
||||
# MHATokenToKVPoolHost and AsymmetricMHATokenToKVPoolHost.
|
||||
orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls
|
||||
orig_mamba_host_pool = assembler.MambaPoolHost
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return orig_kv_host_pool(*args, **kwargs)
|
||||
def get_mha_host_pool_cls_wrapper(device_pool):
|
||||
host_pool_cls = orig_get_mha_host_pool_cls(device_pool)
|
||||
|
||||
def kv_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
return host_pool_cls(*args, **kwargs)
|
||||
|
||||
return kv_host_pool_wrapper
|
||||
|
||||
def mamba_host_pool_wrapper(*args, **kwargs):
|
||||
kwargs["pin_memory"] = False
|
||||
@@ -2384,8 +2400,8 @@ class UnifiedRadixCacheSuite:
|
||||
patchers = [
|
||||
mock.patch.object(
|
||||
assembler,
|
||||
"MHATokenToKVPoolHost",
|
||||
side_effect=kv_host_pool_wrapper,
|
||||
"get_mha_host_pool_cls",
|
||||
side_effect=get_mha_host_pool_cls_wrapper,
|
||||
),
|
||||
mock.patch.object(
|
||||
assembler,
|
||||
|
||||
Reference in New Issue
Block a user