[minimax-m3] Split 2/4: mem-cache / HiCache / sparse KV pool (#28713)

Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
Xinyuan Tong
2026-06-28 00:19:38 +08:00
committed by GitHub
co-authored by hzh0425
parent cfd911ad6e
commit 592f6c849b
10 changed files with 1715 additions and 39 deletions
+8
View File
@@ -297,6 +297,10 @@ class Envs:
SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(True)
SGLANG_TEST_DISAGG_FAILURE_PROB = EnvFloat(0.0)
# HND KV layout folds (page, head) into one paged index for per-kv-head sparse
# page tables (DP attn); paged backends like trtllm_mha consume it directly.
SGLANG_USE_HND_KVCACHE = EnvBool(False)
# Scheduler: memory leak test
SGLANG_TEST_RETRACT = EnvBool(False)
SGLANG_TEST_RETRACT_INTERVAL = EnvInt(3)
@@ -883,6 +887,10 @@ class Envs:
# MiniMax-M3 sparse decode indexer: single JIT radix-select kernel replaces the 2-stage split-K Triton topk.
SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX = EnvBool(True)
# Fused JIT store (minimax_store_kv_index) of main+index K/V instead of separate
# set_*_buffer copies; falls back when main/index dtypes differ or non-CUDA.
SGLANG_OPT_USE_MINIMAX_FUSED_KV_INDEX_STORE = EnvBool(True)
# MiniMax-M3 MXFP8 MoE experimental fusion toggles (default off; A/B only).
SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False)
SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False)
+25 -2
View File
@@ -44,6 +44,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
from sglang.srt.mem_cache.memory_pool import (
DSATokenToKVPool,
MHATokenToKVPool,
MiniMaxSparseKVPool,
MLATokenToKVPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
@@ -92,6 +93,9 @@ class HiRadixCache(RadixCache):
elif isinstance(self.kv_cache, DSATokenToKVPool):
# Filled by attach_hybrid_dsa_pool_to_hiradix_cache after storage extra_config is parsed.
self.token_to_kv_pool_host = None
elif isinstance(self.kv_cache, MiniMaxSparseKVPool):
# Filled by attach_hybrid_minimax_sparse_pool_to_hiradix_cache.
self.token_to_kv_pool_host = None
elif isinstance(self.kv_cache, MLATokenToKVPool):
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache,
@@ -102,7 +106,7 @@ class HiRadixCache(RadixCache):
allocator_type=server_args.hicache_storage_backend,
)
else:
raise ValueError("HiRadixCache only supports MHA, MLA, and DSA models")
raise ValueError("HiRadixCache only supports MHA, MLA, DSA, and MSA models")
self.tp_group = params.tp_cache_group
self.attn_cp_group = params.attn_cp_cache_group
@@ -140,6 +144,22 @@ class HiRadixCache(RadixCache):
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
)
elif isinstance(self.kv_cache, MiniMaxSparseKVPool):
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
attach_hybrid_minimax_sparse_pool_to_hiradix_cache,
)
attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
self,
params,
server_args,
extra_config=extra_config,
prefetch_threshold=prefetch_threshold,
enable_storage_metrics=self.enable_storage_metrics,
load_cache_event=self.load_cache_event,
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
)
else:
self.cache_controller = HiCacheController(
params.token_to_kv_pool_allocator,
@@ -720,7 +740,10 @@ class HiRadixCache(RadixCache):
def _get_extra_pools(self) -> dict:
if not isinstance(self.cache_controller, HybridCacheController):
return {}
if isinstance(self.kv_cache, DSATokenToKVPool):
if isinstance(self.kv_cache, DSATokenToKVPool) or (
isinstance(self.kv_cache, MiniMaxSparseKVPool)
and self.kv_cache.index_k_pool is not None
):
pool = PoolTransfer(
name=PoolName.INDEXER,
hit_policy=PoolHitPolicy.ALL_PAGES,
@@ -19,6 +19,7 @@ from sglang.srt.mem_cache.memory_pool_host import (
HostPoolGroup,
LogicalHostPool,
MambaPoolHost,
MHATokenToKOnlyPoolHost,
MLATokenToKVPoolHost,
PoolEntry,
get_mha_host_pool_cls,
@@ -963,6 +964,69 @@ class _DsaStrategy(StackStrategy):
)
class _MiniMaxSparseStrategy(StackStrategy):
def matches(self, kvcache, components):
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
return isinstance(kvcache, MiniMaxSparseKVPool) and components == {
ComponentType.FULL
}
def build(
self,
*,
cache,
kvcache,
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
model_name=None,
enable_storage_metrics=False,
):
host_pool_group, cache_controller = build_minimax_sparse_hicache_stack(
params=params,
server_args=server_args,
sparse_pool=kvcache,
page_size=cache.page_size,
tp_group=params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
sidecars = []
pools_desc = "KV"
if kvcache.index_k_pool is not None:
sidecars.append(
SidecarPoolSpec(
pool_name=PoolName.INDEXER,
indices_from_pool=PoolName.KV,
)
)
pools_desc = "KV + INDEXER(k-only)"
return StackBuildResult(
host_pool_group=host_pool_group,
cache_controller=cache_controller,
component_host_pools={
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
},
sidecars=sidecars,
transfer_layer_num=kvcache.main_pool.layer_num,
pools_desc=pools_desc,
)
class _PlainKvStrategy(StackStrategy):
def matches(self, kvcache, components):
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
@@ -971,12 +1035,19 @@ class _PlainKvStrategy(StackStrategy):
from sglang.srt.mem_cache.memory_pool import (
DSATokenToKVPool,
HybridLinearKVPool,
MiniMaxSparseKVPool,
)
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
if isinstance(
kvcache,
(SWAKVPool, HybridLinearKVPool, DSATokenToKVPool, DeepSeekV4TokenToKVPool),
(
SWAKVPool,
HybridLinearKVPool,
DSATokenToKVPool,
MiniMaxSparseKVPool,
DeepSeekV4TokenToKVPool,
),
):
return False
return components == {ComponentType.FULL}
@@ -1037,6 +1108,7 @@ _STRATEGIES: list[StackStrategy] = [
_MambaStrategy(),
_SwaStrategy(),
_DsaStrategy(),
_MiniMaxSparseStrategy(),
_PlainKvStrategy(),
]
@@ -1124,6 +1196,194 @@ def attach_hybrid_pool_to_unified_cache(
raise
def build_minimax_sparse_hicache_stack(
*,
params: CacheInitParams,
server_args: ServerArgs,
sparse_pool: Any,
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
"""KV (main_pool) + INDEXER (index_k_pool) host stack for MiniMax M3 sparse."""
# Mappings are stage-local keyed (controller iterates 0..transfer_layer_num).
# PP>1 stays gated below pending end-to-end validation of the sparse host path.
if pp_size > 1:
raise NotImplementedError(
"MiniMax-M3 sparse HiCache does not support pipeline parallelism "
"(pp_size>1) yet."
)
# mirror HiRadix's guard, which the Unified-tree strategy path otherwise skips.
if sparse_pool.index_kv_pool is not None:
raise ValueError(
"MiniMax sparse HiCache currently supports index-k-only sparse layers; "
"index_kv_pool (value-bearing) layers are not cached/restored yet."
)
main_pool = sparse_pool.main_pool
start_layer = main_pool.start_layer
transfer_layer_num = main_pool.layer_num
# Stage-local keys (0..transfer_layer_num) match the controller's per-layer
# load loop; values index the host pool's local layer buffer.
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
kv_host_pool = build_kv_host_pool(
kv_pool=main_pool,
page_size=page_size,
server_args=server_args,
use_mla=False,
)
entries = [
build_pool_entry(
name=PoolName.KV,
host_pool=kv_host_pool,
device_pool=main_pool,
layer_mapping=full_layer_mapping,
transfer_layer_num=transfer_layer_num,
is_anchor=True,
),
]
index_k_pool = sparse_pool.index_k_pool
if index_k_pool is not None:
index_host_pool = MHATokenToKOnlyPoolHost(
index_k_pool,
kv_host_pool,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
)
entries.append(
build_pool_entry(
name=PoolName.INDEXER,
host_pool=index_host_pool,
device_pool=index_k_pool,
layer_mapping={
gid - start_layer: sub_id
for gid, sub_id in sparse_pool.index_k_layer_id_mapping.items()
},
transfer_layer_num=transfer_layer_num,
)
)
host_pool_group = HostPoolGroup(entries)
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_group=params.pp_cache_group,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
return host_pool_group, cache_controller
def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
radix_cache: HiRadixCache,
params: CacheInitParams,
server_args: ServerArgs,
*,
extra_config: dict,
prefetch_threshold: int,
enable_storage_metrics: bool,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> None:
"""Attach HostPoolGroup (KV + index K) + HybridCacheController for HiRadixCache."""
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
try:
sparse_pool = radix_cache.kv_cache
if not isinstance(sparse_pool, MiniMaxSparseKVPool):
raise TypeError(
f"Expected MiniMaxSparseKVPool, got {type(sparse_pool).__name__}"
)
if sparse_pool.index_kv_pool is not None:
raise ValueError(
"MiniMax M3 HiCache L2 currently supports index-k-only sparse layers "
"(sparse_disable_index_value=1 for all sparse layers). "
"This model has index_kv_pool layers; INDEXER_KV sidecar is not "
"implemented yet."
)
main_pool = sparse_pool.main_pool
if sparse_pool.index_k_pool is None:
host_pool_group, cache_controller = build_kv_only_stack(
params=params,
server_args=server_args,
kv_pool=main_pool,
full_layer_mapping={
layer_id: layer_id for layer_id in range(main_pool.layer_num)
},
page_size=radix_cache.page_size,
tp_group=radix_cache.tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=False,
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
storage_backend_extra_config=extra_config,
pp_rank=radix_cache.pp_rank,
pp_size=radix_cache.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
pools_desc = "KV"
else:
host_pool_group, cache_controller = build_minimax_sparse_hicache_stack(
params=params,
server_args=server_args,
sparse_pool=sparse_pool,
page_size=radix_cache.page_size,
tp_group=radix_cache.tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=server_args.hicache_storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
storage_backend_extra_config=extra_config,
pp_rank=radix_cache.pp_rank,
pp_size=radix_cache.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
pools_desc = "KV + INDEXER(k-only)"
sparse_pool.register_layer_transfer_counter(cache_controller.layer_done_counter)
radix_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
radix_cache.token_to_kv_pool_host = host_pool_group
radix_cache.cache_controller = cache_controller
logger.info(
"Attached hybrid MiniMax sparse pool stack to HiRadixCache: pools=%s, "
"transfer_layer_num=%s, sparse_index_k_layers=%s",
pools_desc,
main_pool.layer_num,
len(sparse_pool.index_k_layer_id_mapping),
)
except Exception:
logger.exception("attach_hybrid_minimax_sparse_pool_to_hiradix_cache failed")
raise
def attach_hybrid_dsa_pool_to_hiradix_cache(
radix_cache: HiRadixCache,
params: CacheInitParams,
+516 -35
View File
@@ -1243,39 +1243,42 @@ class MHATokenToKVPool(KVCache):
else v_head_dim if v_head_dim is not None else head_dim
)
# Optional SHUFFLE 5D ("vectorized") physical layout for K/V.
# Selected by `SGLANG_AITER_KV_CACHE_LAYOUT=vectorized_5d` on the ROCm
# AITER backend (HIP + SGLANG_USE_AITER=1). When active:
# K shape: (num_blocks, H, D_k // X, page, X)
# V shape: (num_blocks, H, page // X, D_v, X) where X = 16 / dtype_bytes
# aiter `mha_batch_prefill_func` consumes these 5D shapes natively and
# aiter `pa_decode_gluon` reads SHUFFLE blocks directly during decode.
# An explicit `kv_cache_layout=` argument always wins (e.g. SWAKVPool
# passes "nhd" to keep its SWA sub-pool on the legacy layout); on
# non-AITER platforms the env var is ignored and NHD is forced since
# no consumer kernel exists for SHUFFLE 5D outside the AITER backend.
self.kv_cache_layout = "nhd"
if _use_aiter:
layout = envs.SGLANG_AITER_KV_CACHE_LAYOUT.get().lower()
if layout not in ("nhd", "vectorized_5d"):
raise ValueError(
f"Unsupported SGLANG_AITER_KV_CACHE_LAYOUT={layout!r}; "
"expected 'nhd' or 'vectorized_5d'."
)
self.kv_cache_layout = layout
if layout == "vectorized_5d":
# X is the inner vectorization width in the SHUFFLE layout,
# determined by the STORAGE dtype (not the compute dtype) since
# it controls how many elements fit in 16 bytes of the on-pool
# tensor. For fp8 storage X=16, for bf16/fp16 X=8.
self._kv_vector_x = 16 // self.store_dtype.itemsize
assert (self.size + self.page_size) % self.page_size == 0
assert self.page_size % self._kv_vector_x == 0, (
f"page_size={self.page_size} must be divisible by "
f"X={self._kv_vector_x} for vectorized_5d layout"
)
assert self.head_dim % self._kv_vector_x == 0
assert self.v_head_dim % self._kv_vector_x == 0
# Layout: NHD (default) | HND (SGLANG_USE_HND_KVCACHE) | vectorized_5d (ROCm AITER).
# HND folds (page, head) into one paged index for per-kv-head sparse page tables
# (paged backends like trtllm_mha consume directly). vectorized_5d SHUFFLE 5D:
# K: (num_blocks, H, D_k // X, page, X) V: (num_blocks, H, page // X, D_v, X),
# X = 16 / dtype_bytes — AITER-only (ignored elsewhere, no consumer kernel).
# HND and vectorized_5d are mutually exclusive; HND takes precedence.
self.use_hnd = envs.SGLANG_USE_HND_KVCACHE.get()
if self.use_hnd:
total_slots = self.size + self.page_size
assert total_slots % self.page_size == 0, (
f"HND KV cache needs (size+page_size) divisible by page_size, got "
f"size={self.size}, page_size={self.page_size}"
)
self.num_pages = total_slots // self.page_size
self.kv_cache_layout = "hnd"
else:
self.kv_cache_layout = "nhd"
if _use_aiter:
layout = envs.SGLANG_AITER_KV_CACHE_LAYOUT.get().lower()
if layout not in ("nhd", "vectorized_5d"):
raise ValueError(
f"Unsupported SGLANG_AITER_KV_CACHE_LAYOUT={layout!r}; "
"expected 'nhd' or 'vectorized_5d'."
)
self.kv_cache_layout = layout
if layout == "vectorized_5d":
# X = 16 / storage itemsize: sized by the STORAGE dtype (not compute
# dtype) since it tiles the 16-byte on-pool vector.
self._kv_vector_x = 16 // self.store_dtype.itemsize
assert (self.size + self.page_size) % self.page_size == 0
assert self.page_size % self._kv_vector_x == 0, (
f"page_size={self.page_size} must be divisible by "
f"X={self._kv_vector_x} for vectorized_5d layout"
)
assert self.head_dim % self._kv_vector_x == 0
assert self.v_head_dim % self._kv_vector_x == 0
self._create_buffers()
@@ -1288,7 +1291,9 @@ class MHATokenToKVPool(KVCache):
else None
)
if enable_kv_cache_copy:
if enable_kv_cache_copy and not self.use_hnd:
# The tiled byte copy assumes NHD slot-rows; HND uses a (page, off)
# gather in move_kv_cache instead, so skip the slot-row copy config.
self._init_kv_copy_and_warmup()
else:
self._kv_copy_config = None
@@ -1358,7 +1363,29 @@ class MHATokenToKVPool(KVCache):
if self.enable_custom_mem_pool
else nullcontext()
):
if self.kv_cache_layout == "vectorized_5d":
# The padded page (slot 0's page) absorbs dummy padded-token writes.
if self.use_hnd:
k_shape = (
self.num_pages,
self.head_num,
self.page_size,
self.head_dim,
)
v_shape = (
self.num_pages,
self.head_num,
self.page_size,
self.v_head_dim,
)
self.k_buffer = [
torch.zeros(k_shape, dtype=self.store_dtype, device=self.device)
for _ in range(self.layer_num)
]
self.v_buffer = [
torch.zeros(v_shape, dtype=self.store_dtype, device=self.device)
for _ in range(self.layer_num)
]
elif self.kv_cache_layout == "vectorized_5d":
total_slots = self.size + self.page_size
num_blocks = total_slots // self.page_size
x = self._kv_vector_x
@@ -1452,6 +1479,10 @@ class MHATokenToKVPool(KVCache):
# for disagg
def get_contiguous_buf_infos(self):
assert not self.use_hnd, (
"PD-disaggregation KV transfer assumes NHD slot-row layout; "
"HND KV cache (SGLANG_USE_HND_KVCACHE) is not supported with disagg yet."
)
# layer_num x [seq_len, head_num, head_dim]
# layer_num x [page_num, page_size, head_num, head_dim]
kv_data_ptrs = [
@@ -1478,6 +1509,10 @@ class MHATokenToKVPool(KVCache):
return kv_data_ptrs, kv_data_lens, kv_item_lens
def get_cpu_copy(self, indices, mamba_indices=None):
assert not self.use_hnd, (
"CPU KV offload indexes by slot (NHD); HND KV cache "
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
)
current_platform.synchronize()
kv_cache_cpu = []
chunk_size = self.cpu_offloading_chunk_size
@@ -1496,6 +1531,10 @@ class MHATokenToKVPool(KVCache):
return kv_cache_cpu
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
assert not self.use_hnd, (
"CPU KV offload indexes by slot (NHD); HND KV cache "
"(SGLANG_USE_HND_KVCACHE) is not supported with CPU offload yet."
)
current_platform.synchronize()
chunk_size = self.cpu_offloading_chunk_size
for layer_id in range(self.layer_num):
@@ -1591,6 +1630,16 @@ class MHATokenToKVPool(KVCache):
)
return
if self.use_hnd:
# A slot is [page, :, off, :] (not a contiguous row), so scatter by (page, off).
k_buf = self.k_buffer[layer_id - self.start_layer]
v_buf = self.v_buffer[layer_id - self.start_layer]
pages = loc // self.page_size
offs = loc % self.page_size
k_buf[pages, :, offs, :] = cache_k
v_buf[pages, :, offs, :] = cache_v
return
if self.kv_cache_layout == "vectorized_5d":
# Late-import to keep the NHD path import-clean.
from sglang.srt.layers.attention.utils import (
@@ -1724,6 +1773,14 @@ class MHATokenToKVPool(KVCache):
maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc")
maybe_detect_oob(src_loc, 0, size_limit, "move_kv_cache src_loc")
if self.use_hnd:
pages_t, offs_t = tgt_loc // self.page_size, tgt_loc % self.page_size
pages_s, offs_s = src_loc // self.page_size, src_loc % self.page_size
for kb, vb in zip(self.k_buffer, self.v_buffer):
kb[pages_t, :, offs_t, :] = kb[pages_s, :, offs_s, :]
vb[pages_t, :, offs_t, :] = vb[pages_s, :, offs_s, :]
return
if envs.SGLANG_NATIVE_MOVE_KV_CACHE.get():
move_kv_cache_native(self.k_buffer, self.v_buffer, tgt_loc, src_loc)
return
@@ -1792,6 +1849,10 @@ class NoOpMHATokenToKVPool(MHATokenToKVPool):
"""
def _create_buffers(self):
# No-op pool keeps tiny NHD placeholders regardless of SGLANG_USE_HND_KVCACHE
# (no real KV is stored), so force NHD here to keep the store/move fast paths.
self.use_hnd = False
self.kv_cache_layout = "nhd"
# Allocate minimal placeholder buffers. They exist purely so that code
# paths holding `k_buffer` / `v_buffer` references (pointer tables,
# layer-transfer counters, stride arithmetic) keep working without
@@ -2957,3 +3018,423 @@ def masked_set_kv_buffer_kernel(
value = tl.load(v_ptr + pid * v_stride_B + row * v_stride_H + col, mask=mask)
tl.store(v_buffer_ptr + loc * H * D + idx, value, mask=mask)
class MHATokenToKOnlyPool(KVCache):
"""K-only pool for MiniMax sparse layers whose index branch never reads V
(``sparse_disable_index_value``); allocating V would waste memory."""
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,
page_size,
dtype,
layer_num,
device,
enable_memory_saver,
start_layer,
end_layer,
)
self.head_num = head_num
self.head_dim = head_dim
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.enable_custom_mem_pool
else nullcontext()
):
self.k_buffer = [
torch.zeros(
(size + page_size, head_num, head_dim),
dtype=self.store_dtype,
device=device,
)
for _ in range(layer_num)
]
self._finalize_allocation_log(size)
def _get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.k_buffer[layer_id - self.start_layer].view(self.dtype)
return self.k_buffer[layer_id - self.start_layer]
def register_layer_transfer_counter(
self, layer_transfer_counter: LayerDoneCounter
) -> None:
self.layer_transfer_counter = layer_transfer_counter
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
raise NotImplementedError("MHATokenToKOnlyPool does not allocate V")
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
raise NotImplementedError("MHATokenToKOnlyPool does not allocate V")
def set_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
k_scale: Optional[float] = None,
v_scale: Optional[float] = None,
layer_id_override: Optional[int] = None,
) -> None:
# Routed through MiniMaxSparseKVPool.set_index_k_buffer instead.
raise NotImplementedError(
"MHATokenToKOnlyPool: use set_index_k_buffer on the parent "
"MiniMaxSparseKVPool — this pool does not store V"
)
def get_kv_size_bytes(self):
k_size_bytes = sum(get_tensor_size_bytes(k) for k in self.k_buffer)
return k_size_bytes, 0
class MiniMaxSparseKVPool(KVCache):
def __init__(
self,
size: int,
page_size: int,
dtype: torch.dtype,
head_num: int,
head_dim: int,
idx_head_dim: int,
dense_layer_ids: List[int],
sparse_layer_ids: List[int],
device: str,
disable_value_sparse_layer_ids: Optional[List[int]] = None,
enable_memory_saver: bool = False,
index_dtype: Optional[torch.dtype] = None,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
):
# Do not call super().__init__() — delegate to sub-pools instead.
self.size = size
self.page_size = page_size
self.dtype = dtype
self.device = device
local_dense_layer_ids = [
lid for lid in dense_layer_ids if start_layer <= lid < end_layer
]
local_sparse_layer_ids = [
lid for lid in sparse_layer_ids if start_layer <= lid < end_layer
]
index_dtype = index_dtype if index_dtype is not None else dtype
# Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs
# k_only_sparse (index_k_pool holds only K; V is never read).
disable_set = set(disable_value_sparse_layer_ids or [])
local_kv_sparse_layer_ids = [
g for g in local_sparse_layer_ids if g not in disable_set
]
local_k_only_sparse_layer_ids = [
g for g in local_sparse_layer_ids if g in disable_set
]
# Membership check across all sparse layers, regardless of split.
self.sparse_layer_id_mapping: dict[int, int] = {
gid: i for i, gid in enumerate(local_sparse_layer_ids)
}
# Per-sub-pool local indices.
self.index_kv_layer_id_mapping: dict[int, int] = {
gid: i for i, gid in enumerate(local_kv_sparse_layer_ids)
}
self.index_k_layer_id_mapping: dict[int, int] = {
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
}
self.main_pool = MHATokenToKVPool(
size=size,
page_size=page_size,
dtype=dtype,
head_num=head_num,
head_dim=head_dim,
layer_num=len(local_dense_layer_ids) + len(local_sparse_layer_ids),
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=end_layer,
)
self.index_kv_pool: Optional[MHATokenToKVPool] = (
MHATokenToKVPool(
size=size,
page_size=page_size,
dtype=index_dtype,
head_num=1,
head_dim=idx_head_dim,
layer_num=len(local_kv_sparse_layer_ids),
device=device,
enable_memory_saver=enable_memory_saver,
)
if local_kv_sparse_layer_ids
else None
)
self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
MHATokenToKOnlyPool(
size=size,
page_size=page_size,
dtype=index_dtype,
head_num=1,
head_dim=idx_head_dim,
layer_num=len(local_k_only_sparse_layer_ids),
device=device,
enable_memory_saver=enable_memory_saver,
)
if local_k_only_sparse_layer_ids
else None
)
self.mem_usage = self.main_pool.mem_usage
if self.index_kv_pool is not None:
self.mem_usage += self.index_kv_pool.mem_usage
if self.index_k_pool is not None:
self.mem_usage += self.index_k_pool.mem_usage
# HiCacheController reads these from the top-level KV pool wrapper.
self.layer_num = self.main_pool.layer_num
self.start_layer = self.main_pool.start_layer
self.end_layer = self.main_pool.end_layer
# PD disaggregation reads these directly (no fallback) off the wrapper.
self.head_num = self.main_pool.head_num
self.head_dim = self.main_pool.head_dim
self.layer_transfer_counter = None
def register_layer_transfer_counter(
self, layer_transfer_counter: LayerDoneCounter
) -> None:
self.layer_transfer_counter = layer_transfer_counter
def _wait_for_layer(self, layer_id: int) -> None:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
self._wait_for_layer(layer_id)
return self.main_pool.get_key_buffer(layer_id)
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
self._wait_for_layer(layer_id)
return self.main_pool.get_value_buffer(layer_id)
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
self._wait_for_layer(layer_id)
return self.main_pool.get_kv_buffer(layer_id)
def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
self._wait_for_layer(layer_id)
mapped_id = self.index_kv_layer_id_mapping.get(layer_id)
if mapped_id is None:
raise ValueError(
f"layer_id={layer_id} does not have an index V cache "
f"(either dense, or in the K-only group). "
f"index_kv layers: {list(self.index_kv_layer_id_mapping.keys())}"
)
return self.index_kv_pool.get_kv_buffer(mapped_id)
def get_index_k_buffer(self, layer_id: int) -> torch.Tensor:
self._wait_for_layer(layer_id)
# First try the K-only pool; fall back to the index_kv pool's K side
# so callers that just need K work for both sparse subgroups.
mapped_id = self.index_k_layer_id_mapping.get(layer_id)
if mapped_id is not None:
return self.index_k_pool.get_key_buffer(mapped_id)
mapped_id = self.index_kv_layer_id_mapping.get(layer_id)
if mapped_id is not None:
return self.index_kv_pool.get_key_buffer(mapped_id)
raise ValueError(
f"layer_id={layer_id} is not a sparse attention layer; "
f"sparse layers: {list(self.sparse_layer_id_mapping.keys())}"
)
def set_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
k_scale: float = 1.0,
v_scale: float = 1.0,
) -> None:
"""Write main K/V at `loc`. Works for any layer (dense or sparse)."""
self.main_pool.set_kv_buffer(
layer,
loc,
cache_k,
cache_v,
k_scale,
v_scale,
)
def set_index_kv_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_idx_k: torch.Tensor,
cache_idx_v: torch.Tensor,
k_scale: float = 1.0,
v_scale: float = 1.0,
) -> None:
mapped_id = self.index_kv_layer_id_mapping.get(layer.layer_id)
if mapped_id is None:
raise ValueError(
f"layer.layer_id={layer.layer_id} does not have an index V "
f"cache (either dense, or in the K-only group). "
f"index_kv layers: {list(self.index_kv_layer_id_mapping.keys())}"
)
self.index_kv_pool.set_kv_buffer(
layer,
loc,
cache_idx_k,
cache_idx_v,
k_scale,
v_scale,
layer_id_override=mapped_id,
)
def set_index_k_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_idx_k: torch.Tensor,
) -> None:
mapped_id = self.index_k_layer_id_mapping.get(layer.layer_id)
if mapped_id is None:
raise ValueError(
f"layer.layer_id={layer.layer_id} is not in the K-only "
f"sparse group. K-only layers: "
f"{list(self.index_k_layer_id_mapping.keys())}"
)
sub_pool = self.index_k_pool
if cache_idx_k.dtype != sub_pool.dtype:
cache_idx_k = cache_idx_k.to(sub_pool.dtype)
if sub_pool.store_dtype != sub_pool.dtype:
cache_idx_k = cache_idx_k.view(sub_pool.store_dtype)
sub_pool.k_buffer[mapped_id][loc] = cache_idx_k
def _can_fuse_kv_index_store(
self,
index_pool: MHATokenToKVPool,
cache_k: torch.Tensor,
cache_idx_k: torch.Tensor,
) -> bool:
"""Fast-path precondition: CUDA, no per-store quantization, and a uniform
head byte size shared by main and index caches."""
main = self.main_pool
return (
envs.SGLANG_OPT_USE_MINIMAX_FUSED_KV_INDEX_STORE.get()
and _is_cuda
# No dtype conversion / fp8 scaling on either side (the fused kernel
# is a raw byte copy, it does not quantize).
and main.store_dtype == main.dtype
and index_pool.store_dtype == index_pool.dtype
and cache_k.dtype == main.dtype
and cache_idx_k.dtype == index_pool.dtype
# Uniform head byte size collapses head_dim + dtype into one constant.
and main.dtype == index_pool.dtype
and main.head_dim == index_pool.head_dim
# 128-bit vector copy requires a 16-byte-aligned head size.
and (main.head_dim * main.dtype.itemsize) % 16 == 0
)
def set_fused_kv_index_buffer(
self,
layer: RadixAttention,
loc: torch.Tensor,
cache_k: torch.Tensor,
cache_v: torch.Tensor,
cache_idx_k: torch.Tensor,
cache_idx_v: Optional[torch.Tensor],
) -> None:
"""Store main K/V + index K (+ optional index V) for a sparse layer in
one fused JIT launch, falling back to separate stores when not applicable."""
disable_value = cache_idx_v is None
index_pool = self.index_k_pool if disable_value else self.index_kv_pool
if index_pool is not None and self._can_fuse_kv_index_store(
index_pool, cache_k, cache_idx_k
):
from sglang.jit_kernel.minimax_store_kv_index import store_kv_index
main = self.main_pool
head_bytes = main.head_dim * main.dtype.itemsize
if disable_value:
idx_k_cache = self.get_index_k_buffer(layer.layer_id).flatten(1)
idx_v_cache = None
else:
ik, iv = self.get_index_kv_buffer(layer.layer_id)
idx_k_cache, idx_v_cache = ik.flatten(1), iv.flatten(1)
store_kv_index(
cache_k.flatten(1),
cache_v.flatten(1),
main.get_key_buffer(layer.layer_id).flatten(1),
main.get_value_buffer(layer.layer_id).flatten(1),
cache_idx_k.flatten(1),
idx_k_cache,
None if disable_value else cache_idx_v.flatten(1),
idx_v_cache,
loc,
num_kv_heads=main.head_num,
head_bytes=head_bytes,
)
return
# Fallback: separate stores (identical semantics).
self.set_kv_buffer(layer, loc, cache_k, cache_v)
if disable_value:
self.set_index_k_buffer(layer, loc, cache_idx_k)
else:
self.set_index_kv_buffer(layer, loc, cache_idx_k, cache_idx_v)
def get_kv_size_bytes(self):
sub_pools = [self.main_pool, self.index_kv_pool, self.index_k_pool]
sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None]
return sum(k for k, _ in sizes), sum(v for _, v in sizes)
def get_contiguous_buf_infos(self):
# Main K/V only; index buffers ride the state-buffer channel.
return self.main_pool.get_contiguous_buf_infos()
def get_index_k_state_buf_infos(self):
# Per-page item_len (MHATokenToKVPool convention); index rows share the
# main-KV `loc`, so the transfer reuses the same page-ids.
pool = self.index_k_pool
n = pool.layer_num
data_ptrs = [pool.k_buffer[i].data_ptr() for i in range(n)]
data_lens = [pool.k_buffer[i].nbytes for i in range(n)]
item_lens = [pool.k_buffer[i][0].nbytes * pool.page_size for i in range(n)]
return data_ptrs, data_lens, item_lens
def maybe_get_custom_mem_pool(self):
return self.main_pool.maybe_get_custom_mem_pool()
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
# TODO: spec-decode needs sub-pools built with enable_kv_cache_copy=True,
# then delegate to main_pool/index_pool.move_kv_cache.
raise NotImplementedError(
"move_kv_cache is not yet supported for MiniMaxSparseKVPool: "
"sub-pools must be built with enable_kv_cache_copy=True first."
)
def get_v_head_dim(self):
return self.main_pool.get_value_buffer(0).shape[-1]
@@ -38,6 +38,7 @@ from sglang.jit_kernel.hisparse import transfer_cache_dsv4_mla
from sglang.srt.mem_cache.memory_pool import (
DSATokenToKVPool,
MambaPool,
MHATokenToKOnlyPool,
MHATokenToKVPool,
MLATokenToKVPool,
)
@@ -633,6 +634,316 @@ class MHATokenToKVPoolHost(HostKVCache):
return base_aligned and stride % page_size_bytes == 0
class MHATokenToKOnlyPoolHost(HostKVCache):
"""Host pool for MiniMax sparse index-K buffers (no index V)."""
device_pool: MHATokenToKOnlyPool
def __init__(
self,
device_pool: MHATokenToKOnlyPool,
anchor_host: MHATokenToKVPoolHost,
layout: str,
pin_memory: bool = True,
device: str = "cpu",
allocator_type: str = "default",
):
self.device_pool = device_pool
self.page_size = anchor_host.page_size
self.layout = layout
self.pin_memory = pin_memory
self.device = device
self.allocator = get_allocator_from_storage(allocator_type)
self.dtype = device_pool.store_dtype
self.start_layer = device_pool.start_layer
self.end_layer = device_pool.end_layer
self.head_num = device_pool.head_num
self.head_dim = device_pool.head_dim
self.layer_num = device_pool.layer_num
self.element_dim = self.head_num * self.head_dim
self.token_stride_size = self.element_dim * self.dtype.itemsize
self.layout_dim = self.token_stride_size * self.layer_num
self.size = anchor_host.size
self.page_num = anchor_host.page_num
self.size_per_token = self.get_size_per_token()
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 for MiniMax index-K hierarchical cache. "
f"Requesting {requested_bytes / 1e9:.2f} GB but only have "
f"{available_bytes / 1e9:.2f} GB free."
)
logger.info(
"Allocating %.2f GB host memory for MiniMax sparse index-K (layout=%s).",
requested_bytes / 1e9,
layout,
)
self.init_kv_buffer()
self.lock = threading.RLock()
self.clear()
self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel(
element_size=self.token_stride_size
)
self.k_device_ptrs = torch.tensor(
[x.data_ptr() for x in self.device_pool.k_buffer],
dtype=torch.uint64,
device=self.device_pool.device,
)
if self.layout == "page_first":
transposed = self.k_buffer.transpose(0, 1)
self.k_data_refs = [transposed[i] for i in range(self.layer_num)]
elif self.layout == "layer_first":
self.k_data_refs = [self.k_buffer[i] for i in range(self.layer_num)]
else:
self.k_data_refs = []
self.k_data_ptrs = torch.tensor(
[x.data_ptr() for x in self.k_data_refs],
dtype=torch.uint64,
device=self.device_pool.device,
)
def get_size_per_token(self):
return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize
def get_ksize_per_token(self):
return self.get_size_per_token()
def init_kv_buffer(self):
if self.layout == "layer_first":
dims = (self.layer_num, self.size, self.head_num, self.head_dim)
elif self.layout == "page_first":
dims = (self.size, self.layer_num, self.head_num, self.head_dim)
elif self.layout == "page_first_direct":
dims = (
self.page_num,
self.layer_num,
self.page_size,
self.head_num,
self.head_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
self.k_buffer = alloc_func(
dims,
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
)
def get_hybrid_pool_buffer(self):
return [self.k_buffer]
def load_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
):
if io_backend == "kernel":
if self.layout == "layer_first":
if self.can_use_jit:
jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.k_buffer[layer_id],
cache_src=self.k_buffer[layer_id],
indices_dst=device_indices,
indices_src=host_indices,
element_dim=self.element_dim,
)
else:
transfer_kv_per_layer_mla(
src=self.k_buffer[layer_id],
dst=device_pool.k_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
if self.can_use_jit:
jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.k_buffer[layer_id],
cache_src=self.k_data_refs[layer_id],
indices_dst=device_indices,
indices_src=host_indices,
element_dim=self.element_dim,
)
else:
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.token_stride_size,
src_layout_dim=self.layout_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
if self.layout == "layer_first":
transfer_kv_direct(
src_layers=[self.k_buffer[layer_id]],
dst_layers=[device_pool.k_buffer[layer_id]],
src_indices=host_indices,
dst_indices=device_indices,
page_size=self.page_size,
)
elif self.layout == "page_first_direct":
transfer_kv_per_layer_direct_pf_lf(
src_ptrs=[self.k_buffer],
dst_ptrs=[device_pool.k_buffer[layer_id]],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
page_size=self.page_size,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
if io_backend == "kernel":
if self.layout == "layer_first":
if self.can_use_jit:
for layer_id in range(self.layer_num):
jit_transfer_hicache_one_layer_mla(
cache_dst=self.k_buffer[layer_id],
cache_src=device_pool.k_buffer[layer_id],
indices_dst=host_indices,
indices_src=device_indices,
element_dim=self.element_dim,
)
else:
for layer_id in range(self.layer_num):
transfer_kv_per_layer_mla(
src=device_pool.k_buffer[layer_id],
dst=self.k_buffer[layer_id],
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
if self.can_use_jit:
jit_transfer_hicache_all_layer_mla(
ptr_dst=self.k_data_ptrs,
indices_dst=host_indices,
ptr_src=self.k_device_ptrs,
indices_src=device_indices,
cache_dst_stride_bytes=self.layout_dim,
cache_src_stride_bytes=self.token_stride_size,
element_size=self.element_dim * self.dtype.itemsize,
)
else:
transfer_kv_all_layer_mla_lf_pf(
src_layers=self.k_device_ptrs,
dst=self.k_buffer,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
dst_layout_dim=self.layout_dim,
num_layers=self.layer_num,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
if self.layout == "layer_first":
transfer_kv_direct(
src_layers=device_pool.k_buffer,
dst_layers=self.k_data_refs,
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
elif self.layout == "page_first_direct":
transfer_kv_all_layer_direct_lf_pf(
src_ptrs=device_pool.k_buffer,
dst_ptrs=[self.k_buffer],
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
if self.layout == "layer_first":
data_page = self.k_buffer[:, index : index + self.page_size, :, :]
elif self.layout == "page_first":
data_page = self.k_buffer[index : index + self.page_size, :, :, :]
elif self.layout == "page_first_direct":
real_index = index // self.page_size
data_page = self.k_buffer[real_index : real_index + 1, :, :, :, :]
else:
raise ValueError(f"Unsupported layout: {self.layout}")
if flat:
return data_page.flatten()
return data_page
def get_dummy_flat_data_page(self) -> torch.Tensor:
return torch.zeros(
(self.layer_num, self.page_size, self.head_num, self.head_dim),
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
).flatten()
def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None:
if self.layout == "layer_first":
self.k_buffer[:, index : index + self.page_size, :, :] = data_page.reshape(
self.layer_num, self.page_size, self.head_num, self.head_dim
)
elif self.layout == "page_first":
self.k_buffer[index : index + self.page_size, :, :, :] = data_page.reshape(
self.page_size, self.layer_num, self.head_num, self.head_dim
)
elif self.layout == "page_first_direct":
real_index = index // self.page_size
self.k_buffer[real_index : real_index + 1, :, :, :, :] = data_page.reshape(
1, self.layer_num, self.page_size, self.head_num, self.head_dim
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
def get_page_buffer_meta(self, indices):
"""Meta data for zero-copy storage I/O."""
assert len(indices) % self.page_size == 0
if self.layout not in ["page_first", "page_first_direct"]:
raise ValueError(f"Unsupported layout: {self.layout}")
ptr_list = []
k_buffer_data_ptr = self.k_buffer.data_ptr()
indices = indices.tolist()
for index in range(0, len(indices), self.page_size):
k_ptr = (
k_buffer_data_ptr
+ indices[index]
* self.layer_num
* self.head_num
* self.head_dim
* self.dtype.itemsize
)
ptr_list.append(k_ptr)
element_size = (
self.layer_num
* self.dtype.itemsize
* self.page_size
* self.head_num
* self.head_dim
)
element_size_list = [element_size] * len(ptr_list)
return ptr_list, element_size_list
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.
@@ -8,8 +8,12 @@ import torch
from sglang.srt.configs.model_config import (
get_dsa_index_head_dim,
get_minimax_sparse_attention_config,
get_minimax_sparse_disable_value_layer_ids,
get_minimax_sparse_layer_ids,
is_deepseek_dsa,
is_deepseek_v4,
is_minimax_sparse,
)
from sglang.srt.distributed.parallel_state import (
get_world_group,
@@ -37,6 +41,7 @@ from sglang.srt.mem_cache.memory_pool import (
HybridReqToTokenPool,
MHATokenToKVPool,
MHATokenToKVPoolFP4,
MiniMaxSparseKVPool,
MLATokenToKVPool,
MLATokenToKVPoolFP4,
NoOpMHATokenToKVPool,
@@ -729,6 +734,33 @@ class ModelRunnerKVCacheMixin:
),
**kwargs,
)
elif is_minimax_sparse(self.model_config.hf_config):
_hf_config = self.model_config.hf_config
sparse_cfg = get_minimax_sparse_attention_config(_hf_config)
dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids(
sparse_cfg
)
disable_value_sparse_layer_ids = (
get_minimax_sparse_disable_value_layer_ids(sparse_cfg)
)
self.token_to_kv_pool = MiniMaxSparseKVPool(
size=self.max_total_num_tokens,
page_size=self.page_size,
dtype=self.kv_cache_dtype,
index_dtype=self.dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
),
head_dim=self.model_config.head_dim,
idx_head_dim=sparse_cfg["sparse_index_dim"],
dense_layer_ids=dense_layer_ids,
sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
)
elif config := self.mambaish_config:
extra_args = {}
if self.use_mla_backend:
@@ -21,8 +21,12 @@ import torch
from sglang.srt.configs.model_config import (
get_dsa_index_head_dim,
get_minimax_sparse_attention_config,
get_minimax_sparse_disable_value_layer_ids,
get_minimax_sparse_layer_ids,
is_deepseek_dsa,
is_deepseek_v4,
is_minimax_sparse,
)
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_tp_size
@@ -200,6 +204,44 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
DSATokenToKVPool.index_k_with_scale_buffer_dtype
)
cell_size += indexer_size_per_token * num_layers * element_size
elif is_minimax_sparse(model_config.hf_config):
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
# (sparse-only, single-head; kv layers store K+V, k-only layers store K).
sparse_cfg = get_minimax_sparse_attention_config(model_config.hf_config)
dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids(sparse_cfg)
indexer_k_only_layer_ids = set(
get_minimax_sparse_disable_value_layer_ids(sparse_cfg)
)
local_dense_layer_ids = [
l for l in dense_layer_ids if mr.start_layer <= l < mr.end_layer
]
local_sparse_layer_ids = [
l for l in sparse_layer_ids if mr.start_layer <= l < mr.end_layer
]
num_dense = len(local_dense_layer_ids)
num_sparse = len(local_sparse_layer_ids)
num_indexer_k_only = sum(
1 for l in local_sparse_layer_ids if l in indexer_k_only_layer_ids
)
num_indexer_kv = num_sparse - num_indexer_k_only
kv_heads = model_config.get_num_kv_heads(get_attention_tp_size())
head_dim = model_config.head_dim
indexer_head_dim = sparse_cfg["sparse_index_dim"]
indexer_dtype_size = torch._utils._element_size(mr.dtype)
main_pool_bytes = (
(num_dense + num_sparse) * 2 * kv_heads * head_dim * kv_size
)
indexer_bytes = (
(num_indexer_kv * 2 + num_indexer_k_only)
* indexer_head_dim
* indexer_dtype_size
)
# FP4 scale buffer adjustment doesn't apply to MiniMax sparse:
# cell_size is already a sum over heterogeneous sub-pools.
return main_pool_bytes + indexer_bytes
else:
cell_size = (
model_config.get_num_kv_heads(tp_size)
@@ -0,0 +1,410 @@
import unittest
import psutil
import torch
from sglang.srt.mem_cache.hicache_storage import PoolHitPolicy, PoolName
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
HybridCacheController,
)
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
from sglang.srt.mem_cache.memory_pool_host import (
HICACHE_HOST_MEMORY_RESERVE_BYTES,
MHATokenToKOnlyPoolHost,
MHATokenToKVPoolHost,
)
from sglang.srt.mem_cache.pool_host.common import (
ALLOC_MEMORY_FUNCS,
alloc_with_pin_memory,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-small")
def _cuda_major() -> int:
cuda = getattr(torch.version, "cuda", None)
try:
return int(cuda.split(".")[0]) if cuda else 0
except ValueError:
return 0
# direct+page_first_direct routes to transfer_kv_all_layer_direct_lf_pf, which on
# CUDA 13 throws (cudaErrorInvalidValue) instead of falling back. M3 uses kernel+layer_first.
_DIRECT_PF_BATCHCOPY_BROKEN_CUDA13 = _cuda_major() >= 13
class _FakeLayerTransferCounter:
def __init__(self):
self.waited_layers = []
def wait_until(self, layer_id: int):
self.waited_layers.append(layer_id)
def _make_cpu_minimax_sparse_pool(start_layer: int = 4) -> MiniMaxSparseKVPool:
end_layer = start_layer + 4
return MiniMaxSparseKVPool(
size=8,
page_size=4,
dtype=torch.float32,
head_num=1,
head_dim=2,
idx_head_dim=3,
dense_layer_ids=[start_layer, start_layer + 2],
sparse_layer_ids=[start_layer + 1, start_layer + 3],
disable_value_sparse_layer_ids=[start_layer + 1, start_layer + 3],
device="cpu",
start_layer=start_layer,
end_layer=end_layer,
)
class TestMiniMaxSparseHiCacheIntegration(unittest.TestCase):
def test_hiradix_extra_pools_include_minimax_indexer(self):
pool = _make_cpu_minimax_sparse_pool()
cache = object.__new__(HiRadixCache)
cache.cache_controller = object.__new__(HybridCacheController)
cache.kv_cache = pool
extra = HiRadixCache._get_extra_pools(cache)
transfers = extra["extra_pools"]
self.assertEqual(len(transfers), 1)
self.assertEqual(transfers[0].name, PoolName.INDEXER)
self.assertEqual(transfers[0].indices_from_pool, PoolName.KV)
self.assertEqual(transfers[0].hit_policy, PoolHitPolicy.ALL_PAGES)
def test_index_k_waits_for_full_local_layer(self):
pool = _make_cpu_minimax_sparse_pool()
counter = _FakeLayerTransferCounter()
pool.register_layer_transfer_counter(counter)
pool.get_index_k_buffer(7)
self.assertEqual(counter.waited_layers, [3])
self.assertIsNone(pool.main_pool.layer_transfer_counter)
self.assertIsNone(pool.index_k_pool.layer_transfer_counter)
def test_main_kv_waits_on_minimax_wrapper(self):
pool = _make_cpu_minimax_sparse_pool()
counter = _FakeLayerTransferCounter()
pool.register_layer_transfer_counter(counter)
pool.get_kv_buffer(6)
self.assertEqual(counter.waited_layers, [2])
self.assertIsNone(pool.main_pool.layer_transfer_counter)
self.assertIsNone(pool.index_k_pool.layer_transfer_counter)
def test_k_only_host_pool_layout_contracts(self):
if psutil.virtual_memory().available <= HICACHE_HOST_MEMORY_RESERVE_BYTES:
self.skipTest("Not enough spare host memory for HiCache host pool tests.")
for layout in ("layer_first", "page_first", "page_first_direct"):
with self.subTest(layout=layout):
pool = _make_cpu_minimax_sparse_pool(start_layer=0)
kv_host = MHATokenToKVPoolHost(
device_pool=pool.main_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=pool.page_size,
layout=layout,
pin_memory=False,
device="cpu",
allocator_type="default",
)
index_host = MHATokenToKOnlyPoolHost(
pool.index_k_pool,
kv_host,
layout=layout,
pin_memory=False,
device="cpu",
allocator_type="default",
)
if layout == "layer_first":
self.assertEqual(
index_host.k_buffer.shape,
(
index_host.layer_num,
index_host.size,
index_host.head_num,
index_host.head_dim,
),
)
elif layout == "page_first":
self.assertEqual(
index_host.k_buffer.shape,
(
index_host.size,
index_host.layer_num,
index_host.head_num,
index_host.head_dim,
),
)
else:
self.assertEqual(
index_host.k_buffer.shape,
(
index_host.page_num,
index_host.layer_num,
index_host.page_size,
index_host.head_num,
index_host.head_dim,
),
)
page_start = pool.page_size
flat_page = torch.arange(
index_host.layer_num
* index_host.page_size
* index_host.head_num
* index_host.head_dim,
dtype=index_host.dtype,
)
index_host.set_from_flat_data_page(page_start, flat_page)
self.assertTrue(
torch.equal(index_host.get_data_page(page_start), flat_page)
)
self.assertEqual(
index_host.get_dummy_flat_data_page().numel(), flat_page.numel()
)
self.assertIs(
index_host.get_hybrid_pool_buffer()[0], index_host.k_buffer
)
indices = torch.arange(
page_start,
page_start + pool.page_size,
dtype=torch.int64,
)
if layout == "layer_first":
with self.assertRaisesRegex(ValueError, "layer_first"):
index_host.get_page_buffer_meta(indices)
continue
ptrs, sizes = index_host.get_page_buffer_meta(indices)
self.assertEqual(len(ptrs), 1)
expected_size = (
index_host.layer_num
* index_host.page_size
* index_host.head_num
* index_host.head_dim
* index_host.dtype.itemsize
)
self.assertEqual(sizes, [expected_size] * len(ptrs))
class TestMiniMaxSparseHiCacheTransfer(unittest.TestCase):
def setUp(self):
if not torch.cuda.is_available():
self.skipTest("CUDA is required for MiniMax sparse host transfer tests.")
if is_npu() or is_xpu():
self.skipTest("MiniMax sparse host transfer tests only support CUDA/ROCm.")
if not (is_cuda() or is_hip()):
self.skipTest("CUDA/ROCm not available.")
@staticmethod
def _token_indices_for_pages(pages: torch.Tensor, page_size: int, device: str):
parts = [
torch.arange(
int(page_id) * page_size,
(int(page_id) + 1) * page_size,
device=device,
dtype=torch.int64,
)
for page_id in pages.tolist()
]
return torch.cat(parts, dim=0)
@staticmethod
def _host_k_page(host_pool, layer_id: int, page_id: int, page_size: int):
start = page_id * page_size
if host_pool.layout == "layer_first":
return host_pool.k_buffer[layer_id][start : start + page_size]
if host_pool.layout == "page_first":
return host_pool.k_buffer[start : start + page_size, layer_id]
if host_pool.layout == "page_first_direct":
return host_pool.k_buffer[page_id, layer_id]
raise ValueError(f"Unsupported layout: {host_pool.layout}")
@staticmethod
def _host_v_page(host_pool, layer_id: int, page_id: int, page_size: int):
start = page_id * page_size
if host_pool.layout == "layer_first":
return host_pool.v_buffer[layer_id][start : start + page_size]
if host_pool.layout == "page_first":
return host_pool.v_buffer[start : start + page_size, layer_id]
if host_pool.layout == "page_first_direct":
return host_pool.v_buffer[page_id, layer_id]
raise ValueError(f"Unsupported layout: {host_pool.layout}")
def _run_device_to_host_copy(self, io_backend: str, layout: str):
page_size = 64
layer_num = 4
size = page_size * 4
dense_layer_ids = [0, 1]
sparse_layer_ids = [2, 3]
device_pool = MiniMaxSparseKVPool(
size=size,
page_size=page_size,
dtype=torch.bfloat16,
head_num=4,
head_dim=64,
idx_head_dim=128,
dense_layer_ids=dense_layer_ids,
sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=sparse_layer_ids,
device="cuda",
start_layer=0,
end_layer=layer_num,
)
assert device_pool.index_kv_pool is None
assert device_pool.index_k_pool is not None
pin_memory = io_backend == "kernel"
original_alloc = ALLOC_MEMORY_FUNCS["cuda"]
if pin_memory:
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
try:
kv_host = MHATokenToKVPoolHost(
device_pool=device_pool.main_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout=layout,
pin_memory=pin_memory,
device="cpu",
allocator_type="default",
)
index_host = MHATokenToKOnlyPoolHost(
device_pool.index_k_pool,
kv_host,
layout=layout,
pin_memory=pin_memory,
device="cpu",
allocator_type="default",
)
finally:
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
for layer_id in range(layer_num):
k_main, v_main = device_pool.get_kv_buffer(layer_id)
k_main.copy_(torch.randn_like(k_main) + float(layer_id))
v_main.copy_(torch.randn_like(v_main) + float(layer_id))
for local_id, global_id in enumerate(sparse_layer_ids):
idx_k = device_pool.index_k_pool.k_buffer[local_id]
idx_k.copy_(torch.randn_like(idx_k) + float(global_id) + 100.0)
device_pages = torch.tensor([1, 2, 3], device="cuda", dtype=torch.int64)
host_pages = torch.tensor(
[0, 1, 2],
device="cuda" if io_backend == "kernel" else "cpu",
dtype=torch.int64,
)
device_indices = self._token_indices_for_pages(
device_pages, page_size, device="cuda"
)
host_indices = self._token_indices_for_pages(
host_pages,
page_size,
device="cuda" if io_backend == "kernel" else "cpu",
)
# page_first main-KV backup (staged_write_back.cuh) needs CPU dst_indices,
# index-k backup (hicache.cuh) needs CUDA indices — feed a CPU copy to main only.
kv_host_indices = (
host_indices.cpu()
if (io_backend, layout) == ("kernel", "page_first")
else host_indices
)
kv_host.backup_from_device_all_layer(
device_pool.main_pool, kv_host_indices, device_indices, io_backend
)
index_host.backup_from_device_all_layer(
device_pool.index_k_pool, host_indices, device_indices, io_backend
)
for layer_id in range(layer_num):
for host_page, device_page in zip(
host_pages.tolist(), device_pages.tolist()
):
device_start = device_page * page_size
got_k = self._host_k_page(kv_host, layer_id, host_page, page_size).cpu()
expected_k = device_pool.main_pool.k_buffer[layer_id][
device_start : device_start + page_size
].cpu()
self.assertTrue(torch.equal(got_k, expected_k))
got_v = self._host_v_page(kv_host, layer_id, host_page, page_size).cpu()
expected_v = device_pool.main_pool.v_buffer[layer_id][
device_start : device_start + page_size
].cpu()
self.assertTrue(torch.equal(got_v, expected_v))
for local_id, global_id in enumerate(sparse_layer_ids):
for host_page, device_page in zip(
host_pages.tolist(), device_pages.tolist()
):
got = self._host_k_page(
index_host, local_id, host_page, page_size
).cpu()
expected = device_pool.index_k_pool.k_buffer[local_id][
device_page * page_size : (device_page + 1) * page_size
].cpu()
self.assertTrue(torch.equal(got, expected))
# Round-trip H2D for one sparse index layer.
reload_pages = torch.tensor([0, 1], device="cuda", dtype=torch.int64)
host_device = "cuda" if io_backend == "kernel" else "cpu"
reload_host_pages = torch.tensor([3, 0], device=host_device, dtype=torch.int64)
reload_device_indices = self._token_indices_for_pages(
reload_pages, page_size, device="cuda"
)
reload_host_indices = self._token_indices_for_pages(
reload_host_pages, page_size, device=host_device
)
device_pool.index_k_pool.k_buffer[0].zero_()
index_host.load_to_device_per_layer(
device_pool.index_k_pool,
reload_host_indices,
reload_device_indices,
0,
io_backend,
)
for host_page, device_page in zip(
reload_host_pages.tolist(), reload_pages.tolist()
):
got = device_pool.index_k_pool.k_buffer[0][
device_page * page_size : (device_page + 1) * page_size
].cpu()
expected = self._host_k_page(index_host, 0, host_page, page_size).cpu()
self.assertTrue(torch.equal(got, expected))
def test_device_to_host_kernel_layer_first(self):
self._run_device_to_host_copy(io_backend="kernel", layout="layer_first")
def test_device_to_host_kernel_page_first(self):
self._run_device_to_host_copy(io_backend="kernel", layout="page_first")
def test_device_to_host_direct_layer_first(self):
self._run_device_to_host_copy(io_backend="direct", layout="layer_first")
@unittest.skipIf(
_DIRECT_PF_BATCHCOPY_BROKEN_CUDA13,
"direct+page_first_direct host transfer hits cudaMemcpyBatchAsync "
"cudaErrorInvalidValue on CUDA 13 (sgl-kernel transfer_kv_all_layer_direct_lf_pf "
"throws instead of falling back to per-page copy); M3 production uses "
"io_backend=kernel + layer_first, not this combo.",
)
def test_device_to_host_direct_page_first_direct(self):
self._run_device_to_host_copy(io_backend="direct", layout="page_first_direct")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,58 @@
import unittest
import torch
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
"""Mirror the released MiniMax-M3 config shape: all sparse layers K-only."""
dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2]
sparse_layer_ids = [start_layer + 3 + i for i in range(4)]
end_layer = sparse_layer_ids[-1] + 1
return MiniMaxSparseKVPool(
size=8,
page_size=4,
dtype=torch.float32,
head_num=2,
head_dim=8,
idx_head_dim=16,
dense_layer_ids=dense_layer_ids,
sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=sparse_layer_ids,
device="cpu",
start_layer=start_layer,
end_layer=end_layer,
)
class TestMiniMaxSparsePoolPD(unittest.TestCase):
def test_contiguous_buf_infos_main_only(self):
pool = _make_k_only_pool()
ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
# Main K/V only: 2 entries per main layer (K then V), no index buffers.
n = pool.main_pool.layer_num
self.assertEqual(len(ptrs), 2 * n)
self.assertEqual(len(lens), 2 * n)
self.assertEqual(len(item_lens), 2 * n)
self.assertEqual(ptrs, pool.main_pool.get_contiguous_buf_infos()[0])
def test_index_k_state_buf_infos(self):
pool = _make_k_only_pool()
ptrs, lens, item_lens = pool.get_index_k_state_buf_infos()
n = pool.index_k_pool.layer_num
self.assertEqual(len(ptrs), n)
self.assertEqual(len(lens), n)
self.assertEqual(len(item_lens), n)
for i in range(n):
buf = pool.index_k_pool.k_buffer[i]
self.assertEqual(ptrs[i], buf.data_ptr())
self.assertEqual(lens[i], buf.nbytes)
self.assertEqual(item_lens[i], buf[0].nbytes * pool.page_size)
if __name__ == "__main__":
unittest.main()
@@ -1,5 +1,5 @@
import unittest
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
from sglang.srt.mem_cache.hicache_storage import PoolName, SidecarPoolSpec
from sglang.srt.mem_cache.hybrid_cache import hybrid_pool_assembler
@@ -11,6 +11,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
_DeepSeekV4Strategy,
_DsaStrategy,
_MambaStrategy,
_MiniMaxSparseStrategy,
_PlainKvStrategy,
_select_strategy,
_SwaStrategy,
@@ -36,6 +37,9 @@ class TestUnifiedRadixHiCacheDispatch(unittest.TestCase):
order = [type(s) for s in _STRATEGIES]
# DeepSeekV4 inherits from SWAKVPool, so it must resolve before _SwaStrategy.
self.assertLess(order.index(_DeepSeekV4Strategy), order.index(_SwaStrategy))
self.assertLess(
order.index(_MiniMaxSparseStrategy), order.index(_PlainKvStrategy)
)
self.assertEqual(order[-1], _PlainKvStrategy)
def test_deepseek_v4_full_swa(self):
@@ -68,6 +72,53 @@ class TestUnifiedRadixHiCacheDispatch(unittest.TestCase):
strategy = _select_strategy(kvcache, {FULL})
self.assertIsInstance(strategy, _DsaStrategy)
def test_minimax_sparse(self):
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
kvcache = _mock_kvcache(MiniMaxSparseKVPool)
strategy = _select_strategy(kvcache, {FULL})
self.assertIsInstance(strategy, _MiniMaxSparseStrategy)
def test_minimax_sparse_build_registers_indexer_sidecar(self):
strategy = _MiniMaxSparseStrategy()
host_pool_group = MagicMock()
kv_host_pool = object()
host_pool_group.get_pool.return_value = kv_host_pool
cache_controller = MagicMock()
cache = MagicMock(page_size=4)
kvcache = MagicMock()
kvcache.index_k_pool = object()
kvcache.main_pool.layer_num = 8
params = MagicMock()
params.tp_cache_group = None
params.pp_rank = 0
params.pp_size = 1
server_args = MagicMock()
with patch.object(
hybrid_pool_assembler,
"build_minimax_sparse_hicache_stack",
return_value=(host_pool_group, cache_controller),
) as build_stack:
result = strategy.build(
cache=cache,
kvcache=kvcache,
params=params,
server_args=server_args,
load_cache_event=object(),
)
build_stack.assert_called_once()
self.assertIs(build_stack.call_args.kwargs["sparse_pool"], kvcache)
self.assertIs(result.host_pool_group, host_pool_group)
self.assertIs(result.cache_controller, cache_controller)
self.assertIs(result.component_host_pools[FULL], kv_host_pool)
self.assertEqual(result.pools_desc, "KV + INDEXER(k-only)")
self.assertEqual(result.transfer_layer_num, 8)
self.assertEqual(len(result.sidecars), 1)
self.assertEqual(result.sidecars[0].pool_name, PoolName.INDEXER)
self.assertEqual(result.sidecars[0].indices_from_pool, PoolName.KV)
def test_plain_kv_fallback(self):
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool