[minimax-m3] Split 2/4: mem-cache / HiCache / sparse KV pool (#28713)
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user