diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index c9476a38d..a69a7e6b6 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 65e185413..99d11e1b3 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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, diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index b05523b1b..aa8c9bf14 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -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, diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 6de3e3a1c..83cb33d5b 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -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] diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 00785cf3b..28c97f0cd 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -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. diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index c270db628..2b93a6e58 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -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: diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index 6ac3a0b1d..c5aac6fac 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -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) diff --git a/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py b/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py new file mode 100644 index 000000000..c93336e1a --- /dev/null +++ b/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py @@ -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() diff --git a/test/registered/unit/mem_cache/test_minimax_sparse_pool_pd_unit.py b/test/registered/unit/mem_cache/test_minimax_sparse_pool_pd_unit.py new file mode 100644 index 000000000..71e33e761 --- /dev/null +++ b/test/registered/unit/mem_cache/test_minimax_sparse_pool_pd_unit.py @@ -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() diff --git a/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py index 3ce1ed789..d39d4e564 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py +++ b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py @@ -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