diff --git a/python/sglang/kernels/jit/csrc/minicpm_sala/get_block_table.cuh b/python/sglang/kernels/jit/csrc/minicpm_sala/get_block_table.cuh new file mode 100644 index 000000000..8239fe31a --- /dev/null +++ b/python/sglang/kernels/jit/csrc/minicpm_sala/get_block_table.cuh @@ -0,0 +1,212 @@ +// MiniCPM-SALA sparse attention: build the per-token sparse block table. +// +// Migrated from `3rdparty/sparse_kernel/get_table_kernel.cu`. The original +// CUDA kernels are kept almost verbatim; only the host-side wrappers are +// rewritten from the torch::Tensor + pybind interface to the jit_kernel +// tvm::ffi::TensorView + TensorMatcher/LaunchKernel convention. +// +// The Python wrapper compiles and caches one module per sparse layout. + +#pragma once + +#include // For TensorMatcher, SymbolicSize, SymbolicDevice + +#include // For LaunchKernel + +#include +#include + +#include + +namespace sglang::minicpm_sala { + +constexpr int kTopkPerBlock = 16; + +// topk_idx: [head_group, token_num, kSparseTopK] int32 +// block_table: [batch_size, seqlen_q_max] int32 +// token_to_bs: [token_num] int32 +// token_pos_in_bs: [token_num] int32 +// seqlen_q: [batch_size] int32 +// out_block_table: [token_num, head_group, kSparseTopK * kSparseBlockSize] int32 + +// 1 thread calc 64 element of out_block_table. +// This allows topk_idx to be read once and all corresponding +// out_block_table elements calculated, reducing memory access. +template +__global__ void get_block_table_cuda_blockwise( + const int* topk_idx, + const int* block_table, + const int* token_to_bs, + const int* token_pos_in_bs, + const int* seqlen_q, + int* out_block_table, + const int seqlen_q_max, + const int token_num) { + int token_idx = (blockIdx.x * blockDim.x + threadIdx.x) / (kSparseTopK * kHeadGroup); + if (token_idx >= token_num) return; + int head_group_idx = ((blockIdx.x * blockDim.x + threadIdx.x) / kSparseTopK) % kHeadGroup; + int topk_idx_in_head = (blockIdx.x * blockDim.x + threadIdx.x) % kSparseTopK; + int bs = token_to_bs[token_idx]; + int pos_in_bs = token_pos_in_bs[token_idx]; + int seqlen_q_bs = seqlen_q[bs]; + int sparse_block_idx = + topk_idx[head_group_idx * token_num * kSparseTopK + token_idx * kSparseTopK + topk_idx_in_head]; + + auto out_view = reinterpret_cast(out_block_table); + for (int i = 0; i < kSparseBlockSize; i++) { + int token_idx_in_batch = sparse_block_idx * kSparseBlockSize + i; + if (sparse_block_idx >= 0 && token_idx_in_batch < seqlen_q_bs && token_idx_in_batch < pos_in_bs) { + out_view[token_idx][head_group_idx][topk_idx_in_head][i] = + kHeadGroup * block_table[bs * seqlen_q_max + token_idx_in_batch] + head_group_idx; + } else { + out_view[token_idx][head_group_idx][topk_idx_in_head][i] = 0; + } + } +} + +// 1 thread calculates 1 element of out_block_table. A 1024-thread block +// expands 16 selected blocks in parallel when kSparseBlockSize is 64. +template +__global__ void get_block_table_cuda_elementwise( + const int* topk_idx, + const int* block_table, + const int* token_to_bs, + const int* token_pos_in_bs, + const int* seqlen_q, + int* out_block_table, + const int seqlen_q_max, + const int token_num) { + constexpr int kBlockPerTokenHead = kSparseTopK / kTopkPerBlock; + // calc 16 topk -> 1024 output + __shared__ int topk_idx_share[kTopkPerBlock]; + const int tidx = threadIdx.x; + const int bidx = blockIdx.x; + + if (threadIdx.x < kTopkPerBlock) { + topk_idx_share[tidx] = topk_idx[bidx * kTopkPerBlock + tidx]; + } + + __syncthreads(); + + const int head_group_idx = (bidx / kBlockPerTokenHead) / token_num; + const int token_idx = (bidx / kBlockPerTokenHead) % token_num; + const int topk_idx_in_head = bidx % kBlockPerTokenHead * kTopkPerBlock + tidx / kSparseBlockSize; + + const int sparse_block_idx = topk_idx_share[tidx / kSparseBlockSize]; + + const int token_idx_src = sparse_block_idx * kSparseBlockSize + tidx % kSparseBlockSize; + const int token_idx_dst = token_idx * kHeadGroup * kSparseTopK * kSparseBlockSize + + head_group_idx * kSparseTopK * kSparseBlockSize + topk_idx_in_head * kSparseBlockSize + + tidx % kSparseBlockSize; + if (sparse_block_idx < 0) { + out_block_table[token_idx_dst] = 0; + return; + } + + const int bs = token_to_bs[token_idx]; + const int pos_in_bs = token_pos_in_bs[token_idx]; + const int seqlen_q_bs = seqlen_q[bs]; + + if (token_idx_src < seqlen_q_bs && token_idx_src < pos_in_bs) { + out_block_table[token_idx_dst] = kHeadGroup * block_table[bs * seqlen_q_max + token_idx_src] + head_group_idx; + } else { + out_block_table[token_idx_dst] = 0; + } +} + +// Validate all inputs that are shared across the two kernel variants and +// bind the symbolic dims (token_num / batch_size / seqlen_q_max). The output +// tensor is pre-allocated and fully initialized by the selected kernel. +template +void verify_inputs( + tvm::ffi::TensorView out, + tvm::ffi::TensorView topk_idx, + tvm::ffi::TensorView block_table, + tvm::ffi::TensorView token_to_bs, + tvm::ffi::TensorView token_pos_in_bs, + tvm::ffi::TensorView seqlen_q, + host::SymbolicSize& token_num, + host::SymbolicSize& batch_size, + host::SymbolicSize& seqlen_q_max, + host::SymbolicDevice& device) { + using namespace host; + constexpr int64_t kOutLastDim = static_cast(kSparseTopK) * kSparseBlockSize; + + // topk_idx: [kHeadGroup, token_num, kSparseTopK] + TensorMatcher({static_cast(kHeadGroup), token_num, static_cast(kSparseTopK)}) + .with_dtype() + .with_device(device) + .verify(topk_idx); + // block_table: [batch_size, seqlen_q_max] + TensorMatcher({batch_size, seqlen_q_max}) // + .with_dtype() + .with_device(device) + .verify(block_table); + // token_to_bs / token_pos_in_bs: [token_num] + TensorMatcher({token_num}) // + .with_dtype() + .with_device(device) + .verify(token_to_bs) + .verify(token_pos_in_bs); + // seqlen_q: [batch_size] + TensorMatcher({batch_size}) // + .with_dtype() + .with_device(device) + .verify(seqlen_q); + // out: [token_num, kHeadGroup, kSparseTopK * kSparseBlockSize] + TensorMatcher({token_num, static_cast(kHeadGroup), kOutLastDim}) + .with_dtype() + .with_device(device) + .verify(out); +} + +template +void get_block_table( + tvm::ffi::TensorView out, + tvm::ffi::TensorView topk_idx, + tvm::ffi::TensorView block_table, + tvm::ffi::TensorView token_to_bs, + tvm::ffi::TensorView token_pos_in_bs, + tvm::ffi::TensorView seqlen_q) { + using namespace host; + SymbolicSize token_num{"token_num"}, batch_size{"batch_size"}, seqlen_q_max{"seqlen_q_max"}; + SymbolicDevice device; + device.set_options(); + verify_inputs( + out, topk_idx, block_table, token_to_bs, token_pos_in_bs, seqlen_q, token_num, batch_size, seqlen_q_max, device); + + const int n_token = static_cast(token_num.unwrap()); + const int s_q_max = static_cast(seqlen_q_max.unwrap()); + const DLDevice dev = device.unwrap(); + + constexpr int kThreadsPerBlock = 1024; + constexpr int kElementsPerEntry = kElementwise ? kSparseBlockSize : 1; + const int64_t total = static_cast(n_token) * kHeadGroup * kSparseTopK * kElementsPerEntry; + const int64_t num_blocks = (total + kThreadsPerBlock - 1) / kThreadsPerBlock; + + if constexpr (!kElementwise) { + LaunchKernel(num_blocks, kThreadsPerBlock, dev)( + get_block_table_cuda_blockwise, + static_cast(topk_idx.data_ptr()), + static_cast(block_table.data_ptr()), + static_cast(token_to_bs.data_ptr()), + static_cast(token_pos_in_bs.data_ptr()), + static_cast(seqlen_q.data_ptr()), + static_cast(out.data_ptr()), + s_q_max, + n_token); + } else { + LaunchKernel(num_blocks, kThreadsPerBlock, dev)( + get_block_table_cuda_elementwise, + static_cast(topk_idx.data_ptr()), + static_cast(block_table.data_ptr()), + static_cast(token_to_bs.data_ptr()), + static_cast(token_pos_in_bs.data_ptr()), + static_cast(seqlen_q.data_ptr()), + static_cast(out.data_ptr()), + s_q_max, + n_token); + } +} + +} // namespace sglang::minicpm_sala diff --git a/python/sglang/kernels/jit/minicpm_sala/__init__.py b/python/sglang/kernels/jit/minicpm_sala/__init__.py new file mode 100644 index 000000000..7948ebf32 --- /dev/null +++ b/python/sglang/kernels/jit/minicpm_sala/__init__.py @@ -0,0 +1,3 @@ +from sglang.kernels.jit.minicpm_sala.get_block_table import get_block_table + +__all__ = ["get_block_table"] diff --git a/python/sglang/kernels/jit/minicpm_sala/get_block_table.py b/python/sglang/kernels/jit/minicpm_sala/get_block_table.py new file mode 100644 index 000000000..50a19f76e --- /dev/null +++ b/python/sglang/kernels/jit/minicpm_sala/get_block_table.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_get_block_table_module( + topk: int, head_group_num: int, block_size: int +) -> Module: + """Compile and cache the JIT module for a given sparse topk value. + + One module is built per topk value, replacing the original runtime + ``VALUE_SPLITS_SWITCH(topk, ...)`` dispatch with a compile-time template + argument ``kSparseTopK``. + """ + args = make_cpp_args(topk, head_group_num, block_size) + wrappers = [ + ( + "get_block_table_blockwise", + f"minicpm_sala::get_block_table", + ), + ] + if block_size == 64 and topk % 16 == 0: + wrappers.append( + ( + "get_block_table_elementwise", + f"minicpm_sala::get_block_table", + ) + ) + return load_jit( + f"get_block_table_strategies_topk{topk}_g{head_group_num}_b{block_size}", + *args, + cuda_files=["minicpm_sala/get_block_table.cuh"], + cuda_wrappers=wrappers, + ) + + +def get_block_table( + topk_idx: torch.Tensor, + block_table: torch.Tensor, + token_to_bs: torch.Tensor, + token_pos_in_bs: torch.Tensor, + seqlen_q: torch.Tensor, + head_group_num: int = 2, + block_size: int = 64, + *, + elementwise: bool, +) -> torch.Tensor: + if topk_idx.dim() != 3: + raise RuntimeError( + f"topk_idx must be 3D [head_group, token_num, topk], got shape {tuple(topk_idx.shape)}" + ) + token_num = topk_idx.shape[1] + topk = topk_idx.shape[2] + if topk <= 0 or block_size <= 0: + raise RuntimeError( + f"topk and block_size must be positive, got {topk=} and {block_size=}" + ) + kernel_name = ( + "get_block_table_elementwise" + if elementwise and block_size == 64 and topk % 16 == 0 + else "get_block_table_blockwise" + ) + + out = torch.empty( + (token_num, head_group_num, topk * block_size), + dtype=torch.int32, + device=topk_idx.device, + ) + module = _jit_get_block_table_module(topk, head_group_num, block_size) + getattr(module, kernel_name)( + out, topk_idx, block_table, token_to_bs, token_pos_in_bs, seqlen_q + ) + return out diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 09586209d..f9de19ff1 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1217,6 +1217,54 @@ def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict: return overrides +@_register_for("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM") +def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: + if server_args.enable_dp_attention: + raise ValueError("MiniCPM does not support DP attention") + has_sparse_attention = getattr(hf_config, "has_minicpm_sparse_attention", False) + has_hybrid_attention = has_sparse_attention or getattr( + hf_config, "has_lightning_layers", False + ) + overrides: Dict[str, Any] = {} + if has_hybrid_attention: + if server_args.enable_hierarchical_cache: + raise ValueError("MiniCPM SALA does not support hierarchical cache") + overrides["disable_radix_cache"] = True + if envs.SGLANG_MINICPM_FORCE_DENSE.get(): + dense_backends = { + "minicpm_flashattn": ("fa4" if is_blackwell_supported() else "fa3"), + "minicpm_flashinfer": "flashinfer", + } + for backend_field in ( + "attention_backend", + "prefill_attention_backend", + "decode_attention_backend", + ): + dense_backend = dense_backends.get(getattr(server_args, backend_field)) + if dense_backend is not None: + overrides[backend_field] = dense_backend + elif has_sparse_attention: + uses_sparse_backend = server_args.is_attention_backend_not_set() or any( + backend in ("minicpm_flashattn", "minicpm_flashinfer") + for backend in ( + server_args.attention_backend, + server_args.prefill_attention_backend, + server_args.decode_attention_backend, + ) + ) + if uses_sparse_backend and server_args.disaggregation_mode != "null": + raise ValueError( + "MiniCPM sparse attention does not support PD disaggregation" + ) + if server_args.is_attention_backend_not_set(): + overrides["attention_backend"] = ( + "minicpm_flashinfer" + if is_blackwell_supported() + else "minicpm_flashattn" + ) + return overrides + + @_register_for("MiniCPMV4_6ForConditionalGeneration") def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict: if is_sm100_supported() and server_args.attention_backend is None: diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 3e9b06ce2..53fecdb38 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -36,6 +36,7 @@ from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig from sglang.srt.configs.lfm2_vl import Lfm2VlConfig from sglang.srt.configs.locate_anything import LocateAnythingConfig from sglang.srt.configs.longcat_flash import LongcatFlashConfig +from sglang.srt.configs.minicpm import MiniCPMHybridConfig from sglang.srt.configs.minicpmv4_6 import MiniCPMV4_6Config, MiniCPMV4_6VisionConfig from sglang.srt.configs.minimax_vl import MiniMaxM3VLConfig from sglang.srt.configs.muse_glimmer import ( @@ -113,6 +114,7 @@ __all__ = [ "NemotronH_Nano_Omni_Reasoning_V3_Config", "JetNemotronConfig", "JetVLMConfig", + "MiniCPMHybridConfig", "Step3p5Config", "MiniMaxM3VLConfig", "Step3p7Config", diff --git a/python/sglang/srt/configs/hybrid_arch.py b/python/sglang/srt/configs/hybrid_arch.py index 31c284053..bafd9a5cc 100644 --- a/python/sglang/srt/configs/hybrid_arch.py +++ b/python/sglang/srt/configs/hybrid_arch.py @@ -15,6 +15,7 @@ from sglang.srt.configs import ( Lfm2Config, Lfm2MoeConfig, Lfm2VlConfig, + MiniCPMHybridConfig, NemotronH_Nano_VL_V2_Config, NemotronHConfig, Qwen3_5Config, @@ -42,6 +43,8 @@ def hybrid_lightning_config(model_config: ModelConfig): config = model_config.hf_config if isinstance(config, BailingHybridConfig): return config + if isinstance(config, MiniCPMHybridConfig) and config.has_lightning_layers: + return config return None diff --git a/python/sglang/srt/configs/minicpm.py b/python/sglang/srt/configs/minicpm.py new file mode 100644 index 000000000..be7bb430c --- /dev/null +++ b/python/sglang/srt/configs/minicpm.py @@ -0,0 +1,194 @@ +from transformers import PretrainedConfig + +from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape +from sglang.srt.runtime_context import get_parallel + +_MIXER_TYPE_ALIASES = { + "minicpm4": "minicpm4", + "minicpm": "minicpm4", + "standard": "minicpm4", + "attention": "minicpm4", + "attn": "minicpm4", + "lightning": "lightning-attn", + "lightning_attn": "lightning-attn", + "lightning-attn": "lightning-attn", +} + + +class MiniCPMHybridConfig(PretrainedConfig): + """ + Configuration class for hybrid MiniCPM models. + + This config extends PretrainedConfig to match the pattern used by other + hybrid/linear attention models (Falcon H1, Nemotron H, Kimi Linear, etc.) + and provides cache parameters for the Simple GLA attention mechanism. + """ + + model_type = "minicpm_sala" + + def __init__( + self, + # Base model config fields + vocab_size=150528, + hidden_size=4096, + num_hidden_layers=32, + num_attention_heads=32, + num_key_value_heads=8, + head_dim=None, + hidden_act="silu", + intermediate_size=14336, + initializer_range=0.02, + rms_norm_eps=1e-6, + use_cache=True, + pad_token_id=0, + bos_token_id=1, + eos_token_id=2, + tie_word_embeddings=False, + max_position_embeddings=32768, + rope_theta=10000.0, + rope_scaling=None, + scale_emb=12, + scale_depth=1.4, + dim_model_base=256, + # MiniCPM-specific hybrid config fields + mixer_types=None, + lightning_nh=None, + lightning_nkv=None, + lightning_head_dim=None, + lightning_scale="1/sqrt(d)", + lightning_layerwise_decay=False, + lightning_use_rope=True, + use_output_gate=False, + attention_bias=False, + use_output_norm=False, + qk_norm=True, + attn_use_rope=True, + attn_use_output_gate=False, + sparse_config=None, + **kwargs, + ): + for unused_field in ("minicpm4", "lightning", "sparse_use_nope"): + kwargs.pop(unused_field, None) + + self.vocab_size = vocab_size + self.hidden_size = hidden_size + self.num_hidden_layers = num_hidden_layers + self.num_attention_heads = num_attention_heads + self.num_key_value_heads = num_key_value_heads + self.head_dim = ( + head_dim if head_dim is not None else hidden_size // num_attention_heads + ) + self.max_position_embeddings = max_position_embeddings + self.rope_theta = rope_theta + self.rope_scaling = rope_scaling + self.scale_emb = scale_emb + self.scale_depth = scale_depth + self.dim_model_base = dim_model_base + self.hidden_act = hidden_act + self.intermediate_size = intermediate_size + self.initializer_range = initializer_range + self.rms_norm_eps = rms_norm_eps + self.use_cache = use_cache + # Hybrid config fields + if not mixer_types: + mixer_types = ["minicpm4"] + elif len(mixer_types) > num_hidden_layers: + raise ValueError(f"Invalid number of mixer types: {len(mixer_types)}") + try: + mixer_types = [ + _MIXER_TYPE_ALIASES[mixer_type] for mixer_type in mixer_types + ] + except KeyError as exc: + raise ValueError(f"Unsupported mixer type: {exc.args[0]}") from exc + repeats = (num_hidden_layers + len(mixer_types) - 1) // len(mixer_types) + self.mixer_types = (mixer_types * repeats)[:num_hidden_layers] + self.lightning_nh = ( + lightning_nh if lightning_nh is not None else num_attention_heads + ) + self.lightning_nkv = ( + lightning_nkv if lightning_nkv is not None else num_key_value_heads + ) + self.lightning_head_dim = ( + lightning_head_dim if lightning_head_dim is not None else self.head_dim + ) + if ( + "lightning-attn" in self.mixer_types + and self.lightning_nh != self.lightning_nkv + ): + raise ValueError( + "MiniCPM Lightning attention requires equal query and KV head " + "counts because the seg_la backend does not support GQA: " + f"lightning_nh={self.lightning_nh}, " + f"lightning_nkv={self.lightning_nkv}" + ) + self.lightning_scale = lightning_scale + self.lightning_layerwise_decay = lightning_layerwise_decay + self.lightning_use_rope = lightning_use_rope + self.use_output_gate = use_output_gate + self.attention_bias = attention_bias + self.use_output_norm = use_output_norm + self.qk_norm = qk_norm + self.attn_use_rope = attn_use_rope + self.attn_use_output_gate = attn_use_output_gate + self.sparse_config = sparse_config + + super().__init__( + pad_token_id=pad_token_id, + bos_token_id=bos_token_id, + eos_token_id=eos_token_id, + tie_word_embeddings=tie_word_embeddings, + **kwargs, + ) + + @property + def num_linear_key_value_heads(self) -> int: + return self.lightning_nkv + + @property + def mamba2_cache_params(self): + """Return linear-attention cache parameters for lightning layers.""" + lightning_layer_ids = self.lightning_layer_ids + + if ( + not lightning_layer_ids + or not self.lightning_nkv + or not self.lightning_head_dim + ): + return None + + shape = Mamba2StateShape.create( + tp_world_size=get_parallel().attn_tp_size, + intermediate_size=0, + n_groups=0, + num_heads=self.lightning_nkv, + head_dim=self.lightning_head_dim, + state_size=self.lightning_head_dim, + conv_kernel=1, + ) + + return Mamba2CacheParams(shape=shape, layers=lightning_layer_ids) + + @property + def full_attention_layer_ids(self): + return [ + i + for i, mixer_type in enumerate(self.mixer_types) + if mixer_type == "minicpm4" + ] + + @property + def has_minicpm_sparse_attention(self) -> bool: + """Check if this config has MiniCPM sparse attention layers.""" + return self.sparse_config is not None and any( + mt == "minicpm4" for mt in self.mixer_types + ) + + @property + def has_lightning_layers(self) -> bool: + """Check if this config has lightning attention layers.""" + return any(mt == "lightning-attn" for mt in self.mixer_types) + + @property + def lightning_layer_ids(self) -> list: + """Get the indices of layers with lightning attention.""" + return [i for i, mt in enumerate(self.mixer_types) if mt == "lightning-attn"] diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 57c637ffd..667294da4 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -162,6 +162,7 @@ class DecodeReqToTokenPool: # here: HybridMambaDecodeReqToTokenPool borrows this __init__ while # inheriting ReqToTokenPool.alloc, which bumps it. self.req_generation = torch.zeros(self._alloc_size, dtype=torch.int64) + self._aux_cache: Any = None def write(self, indices, values): self.req_to_token[indices] = values @@ -169,6 +170,20 @@ class DecodeReqToTokenPool: def available_size(self): return len(self.free_slots) + def reset_aux_cache_allocator(self) -> None: + pass + + def schedulable_token_capacity(self, physical_capacity: int) -> int: + return physical_capacity + + def alloc_aux_to_lengths( + self, + *, + req_pool_indices_cpu: torch.Tensor, + target_seq_lens_cpu: torch.Tensor, + ) -> None: + pass + def alloc(self, reqs: List[Req]) -> Optional[List[int]]: # Indices of reqs that already have a req_pool_idx and will reuse # their existing slot (e.g. chunked prefill continuing across chunks). diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index b9be96fc2..ec1365875 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1074,6 +1074,11 @@ class Envs: # =================================================================== # Kernel selection and fused backends # =================================================================== + # MiniCPM sparse attention developer switches + SGLANG_MINICPM_FUSE_TOPK = EnvBool(False) + SGLANG_MINICPM_DENSE_AS_SPARSE = EnvBool(False) + SGLANG_MINICPM_FORCE_DENSE = EnvBool(False) + SGLANG_USE_SGL_FA3_KERNEL = EnvBool(True) # Force every sglang.kernels BaseFusedOp onto one backend (a KernelBackend # value, e.g. "torch" / "torch_compile" / "triton" / "aot"); unset = diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index c560895f7..ef9af670e 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -321,6 +321,20 @@ def attn_backend_wrapper_for_draft_decode(runner: "ModelRunner", backend): return backend +@register_attention_backend("minicpm_flashattn") +def create_minicpm_flashattn_backend(runner): + from sglang.srt.layers.attention.minicpm.backend import MiniCPMSparseBackend + + return MiniCPMSparseBackend(runner, use_flashinfer=False) + + +@register_attention_backend("minicpm_flashinfer") +def create_minicpm_flashinfer_backend(runner): + from sglang.srt.layers.attention.minicpm.backend import MiniCPMSparseBackend + + return MiniCPMSparseBackend(runner, use_flashinfer=True) + + def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBackend"): """ Wrapper for special models like hybrid GDN, so we don't diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index e06bcb449..2f7d0d5d6 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -1170,6 +1170,61 @@ class FlashAttentionBackend(AttentionBackend): self.forward_metadata = metadata + def get_paged_mha_kv_cache( + self, + layer: RadixAttention, + *, + head_group_num: int = 1, + ) -> tuple[torch.Tensor, torch.Tensor]: + key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) + return ( + key_cache.view( + -1, + self.page_size, + layer.tp_k_head_num // head_group_num, + layer.head_dim, + ), + value_cache.view( + -1, + self.page_size, + layer.tp_v_head_num // head_group_num, + layer.v_head_dim, + ), + ) + + def prepare_paged_mha_query( + self, + q: torch.Tensor, + q_rope: Optional[torch.Tensor], + k_rope: Optional[torch.Tensor], + layer: RadixAttention, + *, + logical_batch_size: int, + kv_head_num: int, + is_prefill: bool, + ) -> tuple[ + torch.Tensor, + Optional[torch.Tensor], + Optional[torch.Tensor], + Optional[torch.Tensor], + Optional[torch.Tensor], + ]: + k_descale = v_descale = None + if ( + self.kv_cache_dtype_str != "auto" + and layer.head_dim <= 256 + and not self.kv_cache_is_mxfp8 + and (not is_prefill or self.fa_impl_ver != 4) + ): + if layer.k_scale is not None: + descale_shape = (logical_batch_size, kv_head_num) + k_descale = layer.k_scale.expand(descale_shape) + v_descale = layer.v_scale.expand(descale_shape) + q = q.to(self.kv_cache_dtype) + q_rope = q_rope.to(self.kv_cache_dtype) if q_rope is not None else None + k_rope = k_rope.to(self.kv_cache_dtype) if k_rope is not None else None + return q, q_rope, k_rope, k_descale, v_descale + def forward_extend( self, q: torch.Tensor, @@ -1282,24 +1337,15 @@ class FlashAttentionBackend(AttentionBackend): if is_swa_layer else (-1, -1) ) - fa_k_descale, fa_v_descale = None, None - # only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention - # has corresponding quantization method so that layer.k_scale is not None, - # 3) layer.head_dim <= 256 since fa3 kernel require fp16 and bf16 data type in this case, - # 4) fa_impl_ver != 4 since fa4 does not currently support fp8 queries and keys. - if ( - self.kv_cache_dtype_str != "auto" - and layer.head_dim <= 256 - and self.fa_impl_ver != 4 - and not self.kv_cache_is_mxfp8 - ): - if layer.k_scale is not None: - descale_shape = (forward_batch.batch_size, layer.tp_k_head_num) - fa_k_descale = layer.k_scale.expand(descale_shape) - fa_v_descale = layer.v_scale.expand(descale_shape) - q = q.to(self.kv_cache_dtype) - q_rope = q_rope.to(self.kv_cache_dtype) if q_rope is not None else None - k_rope = k_rope.to(self.kv_cache_dtype) if k_rope is not None else None + q, q_rope, k_rope, fa_k_descale, fa_v_descale = self.prepare_paged_mha_query( + q, + q_rope, + k_rope, + layer, + logical_batch_size=forward_batch.batch_size, + kv_head_num=layer.tp_k_head_num, + is_prefill=True, + ) # Check if we should use local attention use_local_attn = ( self.has_local_attention @@ -1385,13 +1431,8 @@ class FlashAttentionBackend(AttentionBackend): # Use Flash Attention for prefill if not self.use_mla: # Do multi-head attention - key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) - - key_cache = key_cache.view( - -1, self.page_size, layer.tp_k_head_num, layer.head_dim - ) - value_cache = value_cache.view( - -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim + key_cache, value_cache = self.get_paged_mha_kv_cache( + layer, ) if layer.is_cross_attention: page_table = metadata.encoder_page_table @@ -1869,34 +1910,23 @@ class FlashAttentionBackend(AttentionBackend): else None ) - fa_k_descale, fa_v_descale = None, None - # only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention - # has corresponding quantization method so that layer.k_scale is not None, - # 3) layer.head_dim <= 256 since fa3 kernel require fp16 and bf16 data type in this case. - if ( - self.kv_cache_dtype_str != "auto" - and layer.head_dim <= 256 - and not self.kv_cache_is_mxfp8 - ): - if layer.k_scale is not None: - descale_shape = (forward_batch.batch_size, layer.tp_k_head_num) - fa_k_descale = layer.k_scale.expand(descale_shape) - fa_v_descale = layer.v_scale.expand(descale_shape) - q = q.to(self.kv_cache_dtype) - q_rope = q_rope.to(self.kv_cache_dtype) if q_rope is not None else None - k_rope = k_rope.to(self.kv_cache_dtype) if k_rope is not None else None + q, q_rope, k_rope, fa_k_descale, fa_v_descale = self.prepare_paged_mha_query( + q, + q_rope, + k_rope, + layer, + logical_batch_size=forward_batch.batch_size, + kv_head_num=layer.tp_k_head_num, + is_prefill=False, + ) if fa_k_descale is not None: kwargs["k_descale"] = fa_k_descale kwargs["v_descale"] = fa_v_descale if not self.use_mla: # Do multi-head attention - key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id) - key_cache = key_cache.view( - -1, self.page_size, layer.tp_k_head_num, layer.head_dim - ) - value_cache = value_cache.view( - -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim + key_cache, value_cache = self.get_paged_mha_kv_cache( + layer, ) if layer.is_cross_attention: diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 703209f03..0de53ce64 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -1061,6 +1061,23 @@ class FlashInferAttnBackend(AttentionBackend): for i in range(self.num_wrappers) ] + def get_cuda_graph_decode_wrappers( + self, + *, + bs: int, + num_tokens: int, + ) -> list: + wrappers = self.decode_cuda_graph_metadata.get(bs) + if wrappers is None: + self._prepare_cuda_graph_metadata( + bs, + num_tokens, + ForwardMode.DECODE, + spec_info=None, + ) + wrappers = self.decode_cuda_graph_metadata[bs] + return wrappers + def _create_prefill_wrappers(self, bs: int, use_custom_mask: bool = False) -> list: # FlashInfer's prefill wrapper decides mask mode based on whether # `custom_mask_buf` is initialized (not whether a custom mask is provided). diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index e5f867caa..85760ea06 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -74,14 +74,20 @@ class LightningAttentionBackend(MambaAttnBackendBase): if hasattr(model_runner.model_config, "block") else 256 ) - total_num_heads = model_runner.model_config.hf_config.num_attention_heads - num_hidden_layers = model_runner.model_config.hf_config.num_hidden_layers + config = model_runner.model_config.hf_config + total_num_heads = getattr( + config, "num_linear_key_value_heads", config.num_attention_heads + ) + layerwise_decay = getattr(config, "lightning_layerwise_decay", True) + assert total_num_heads % get_parallel().attn_tp_size == 0 + num_hidden_layers = config.num_hidden_layers self.tp_slope = LightningAttentionBackend._build_slope_tensor( - total_num_heads, num_hidden_layers, self.device - ) - self.linear_backend = getattr( - model_runner.model_config.hf_config, "linear_backend", "seg_la" + total_num_heads, + num_hidden_layers, + self.device, + layerwise_decay=layerwise_decay, ) + self.linear_backend = getattr(config, "linear_backend", "seg_la") logger.info( f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}" ) @@ -130,7 +136,10 @@ class LightningAttentionBackend(MambaAttnBackendBase): @staticmethod def _build_slope_tensor( - n_attention_heads: int, num_hidden_layers: int, device="cuda" + n_attention_heads: int, + num_hidden_layers: int, + device="cuda", + layerwise_decay: bool = True, ): def get_slopes(n): def get_slopes_power_of_2(n): @@ -153,7 +162,9 @@ class LightningAttentionBackend(MambaAttnBackendBase): tp_heads = n_attention_heads // get_parallel().attn_tp_size tp_rank = get_parallel().attn_tp_rank - if num_hidden_layers <= 1: + if not layerwise_decay: + slope_rate_list = [slopes] * num_hidden_layers + elif num_hidden_layers <= 1: slope_rate_list = [slopes * (1 + 1e-5)] else: slope_rate_list = [ @@ -285,6 +296,7 @@ class LightningAttentionBackend(MambaAttnBackendBase): cache_indices=intermediate_state_indices, track_lens=track_lens, track_state_indices=track_state_indices, + softmax_scale=layer.scaling, decouple=True, ) return hidden diff --git a/python/sglang/srt/layers/attention/minicpm/__init__.py b/python/sglang/srt/layers/attention/minicpm/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/layers/attention/minicpm/attention_adapter.py b/python/sglang/srt/layers/attention/minicpm/attention_adapter.py new file mode 100644 index 000000000..916251388 --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/attention_adapter.py @@ -0,0 +1,288 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.ops.attention.flash_attention import flash_attn_with_kvcache +from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_triton +from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend +from sglang.srt.layers.attention.minicpm.sparse_utils import ( + MiniCPMSparseMetadata, +) +from sglang.srt.utils import is_flashinfer_available + +if TYPE_CHECKING: + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.model_runner import ModelRunner + + +class MiniCPMFlashAttentionAdapter: + def __init__(self, flash_attn_backend: FlashAttentionBackend): + self.flash_attn_backend = flash_attn_backend + + def prepare_forward( + self, + metadata: MiniCPMSparseMetadata, + *, + is_prefill: bool, + graph: bool, + ) -> None: + pass + + def init_cuda_graph_state(self, max_num_tokens: int) -> None: + pass + + def forward( + self, + q: torch.Tensor, + key_cache: torch.Tensor, + value_cache: torch.Tensor, + metadata: MiniCPMSparseMetadata, + layer: RadixAttention, + *, + is_prefill: bool, + k_descale: Optional[torch.Tensor] = None, + v_descale: Optional[torch.Tensor] = None, + sinks: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + kwargs = {} + if sinks is not None: + kwargs["sinks"] = sinks + if k_descale is not None: + kwargs["k_descale"] = k_descale + kwargs["v_descale"] = v_descale + + return flash_attn_with_kvcache( + q=q, + k_cache=key_cache, + v_cache=value_cache, + page_table=metadata.sparse_page_table, + cache_seqlens=metadata.sparse_cache_seqlens_int32, + cu_seqlens_q=metadata.sparse_cu_seqlens_q, + cu_seqlens_k_new=metadata.sparse_cu_seqlens_k, + max_seqlen_q=( + metadata.sparse_max_seq_len_q + if is_prefill + else metadata.base.max_seq_len_q + ), + softmax_scale=layer.scaling, + causal=True, + window_size=(-1, -1), + softcap=layer.logit_cap, + num_splits=self.flash_attn_backend.num_splits, + ver=self.flash_attn_backend.fa_impl_ver, + **kwargs, + ) + + +class MiniCPMFlashInferAdapter: + def __init__( + self, + model_runner: ModelRunner, + *, + head_group_num: int, + heads_per_group: int, + head_dim: int, + page_size: int, + max_kv_tokens_per_row: int, + ): + if not is_flashinfer_available(): + raise RuntimeError("minicpm_flashinfer requires the flashinfer package.") + + from sglang.srt.layers.attention.flashinfer_backend import ( + FlashInferAttnBackend, + ) + + self.device = model_runner.device + self.head_group_num = head_group_num + self.num_qo_heads = heads_per_group + self.num_kv_heads = 1 + self.head_dim = head_dim + self.page_size = page_size + self.max_kv_tokens_per_row = max_kv_tokens_per_row + self.q_dtype = model_runner.dtype + self.kv_dtype = model_runner.kv_cache_dtype + + max_sparse_bs = model_runner.req_to_token_pool.size * head_group_num + self.kv_indptr = torch.zeros( + max_sparse_bs + 1, + dtype=torch.int32, + device=self.device, + ) + self.kv_indices = torch.zeros( + max_sparse_bs * max_kv_tokens_per_row, + dtype=torch.int32, + device=self.device, + ) + self.kv_last_page_len = torch.ones( + max_sparse_bs, + dtype=torch.int32, + device=self.device, + ) + self.rows = torch.arange( + max_sparse_bs, + dtype=torch.int32, + device=self.device, + ) + self.flashinfer_backend = FlashInferAttnBackend( + model_runner, + skip_prefill=False, + kv_indptr_buf=self.kv_indptr, + kv_last_page_len_buf=self.kv_last_page_len, + ) + self.active_wrapper = None + self.active_kv_indptr = None + self.active_kv_indices = None + self.active_rows = None + self.prefill_planned = False + + def prepare_forward( + self, + metadata: MiniCPMSparseMetadata, + *, + is_prefill: bool, + graph: bool, + ) -> None: + if is_prefill: + self.prefill_planned = False + else: + self._prepare( + metadata, + is_prefill=False, + graph=graph, + ) + + def init_cuda_graph_state(self, max_num_tokens: int) -> None: + self.flashinfer_backend.init_cuda_graph_state( + max_num_tokens, + max_num_tokens, + kv_indices_buf=self.kv_indices, + ) + + def _prepare( + self, + metadata: MiniCPMSparseMetadata, + *, + is_prefill: bool, + graph: bool = False, + ) -> None: + cache_seqlens = metadata.sparse_cache_seqlens_int32 + sparse_bs = cache_seqlens.numel() + if sparse_bs == 0: + self.active_wrapper = None + return + + if is_prefill: + kv_indptr = metadata.sparse_cu_seqlens_k + kv_indices = torch.empty( + metadata.sparse_page_table.numel(), + dtype=torch.int32, + device=self.device, + ) + kv_last_page_len = (cache_seqlens > 0).to(torch.int32) + rows = torch.arange( + sparse_bs, + dtype=torch.int32, + device=self.device, + ) + wrapper = self.flashinfer_backend.prefill_wrappers_paged[0] + wrapper.begin_forward( + metadata.sparse_cu_seqlens_q, + kv_indptr, + kv_indices, + kv_last_page_len, + self.num_qo_heads, + self.num_kv_heads, + self.head_dim, + self.page_size, + causal=True, + q_data_type=self.q_dtype, + kv_data_type=self.kv_dtype, + non_blocking=True, + ) + else: + kv_indptr = self.kv_indptr[: sparse_bs + 1] + kv_indptr[0] = 0 + torch.cumsum(cache_seqlens, dim=0, out=kv_indptr[1:]) + kv_indices = self.kv_indices[: sparse_bs * self.max_kv_tokens_per_row] + kv_last_page_len = self.kv_last_page_len[:sparse_bs] + kv_last_page_len.copy_((cache_seqlens > 0).to(torch.int32)) + rows = self.rows[:sparse_bs] + if graph: + graph_bs = sparse_bs // self.head_group_num + wrapper = self.flashinfer_backend.get_cuda_graph_decode_wrappers( + bs=graph_bs, + num_tokens=sparse_bs, + )[0] + else: + wrapper = self.flashinfer_backend.decode_wrappers[0] + + wrapper.begin_forward( + kv_indptr, + kv_indices, + kv_last_page_len, + self.num_qo_heads, + self.num_kv_heads, + self.head_dim, + self.page_size, + q_data_type=self.q_dtype, + kv_data_type=self.kv_dtype, + non_blocking=True, + ) + + self.active_wrapper = wrapper + self.active_kv_indptr = kv_indptr + self.active_kv_indices = kv_indices + self.active_rows = rows + + def forward( + self, + q: torch.Tensor, + key_cache: torch.Tensor, + value_cache: torch.Tensor, + metadata: MiniCPMSparseMetadata, + layer: RadixAttention, + *, + is_prefill: bool, + k_descale: Optional[torch.Tensor] = None, + v_descale: Optional[torch.Tensor] = None, + sinks: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if sinks is not None: + raise NotImplementedError( + "minicpm_flashinfer does not support attention sinks" + ) + if is_prefill and not self.prefill_planned: + self._prepare(metadata, is_prefill=True) + self.prefill_planned = True + + cache_seqlens = metadata.sparse_cache_seqlens_int32 + sparse_bs = cache_seqlens.numel() + create_flashinfer_kv_indices_triton[(sparse_bs,)]( + metadata.sparse_page_table, + self.active_rows, + cache_seqlens, + self.active_kv_indptr, + None, + self.active_kv_indices, + metadata.sparse_page_table.stride(0), + ) + kwargs = { + "sm_scale": layer.scaling, + "logits_soft_cap": layer.logit_cap or None, + "k_scale": layer.k_scale_float, + "v_scale": layer.v_scale_float, + } + if is_prefill: + return self.active_wrapper.forward( + q, + (key_cache, value_cache), + causal=True, + **kwargs, + ) + return self.active_wrapper.forward( + q, + (key_cache, value_cache), + **kwargs, + ) diff --git a/python/sglang/srt/layers/attention/minicpm/backend.py b/python/sglang/srt/layers/attention/minicpm/backend.py new file mode 100644 index 000000000..657c41f97 --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/backend.py @@ -0,0 +1,1128 @@ +from __future__ import annotations + +from types import SimpleNamespace +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.configs.minicpm import MiniCPMHybridConfig +from sglang.srt.environ import envs +from sglang.srt.layers.attention.base_attn_backend import AttentionBackend +from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionBackend, +) +from sglang.srt.layers.attention.minicpm.attention_adapter import ( + MiniCPMFlashAttentionAdapter, + MiniCPMFlashInferAdapter, +) +from sglang.srt.layers.attention.minicpm.cache import attach_compressed_cache +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.runtime_context import get_parallel +from sglang.srt.utils import is_blackwell_supported, next_power_of_2 + +if TYPE_CHECKING: + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.model_runner import ModelRunner + + +from sglang.kernels.jit.minicpm_sala import get_block_table +from sglang.srt.layers.attention.minicpm.sparse_utils import ( + CompressionLevelMetadata, + MiniCPMSparseMetadata, + _build_k1_k2_compression_metadata, + _plan_sparse_decode, + _plan_sparse_prefill, + allocate_and_compress_keys, + batched_gather, + compressed_attention, + compressed_attention_tilelang, + get_compress_k_v2, +) + + +def _transpose_head_group_layout( + tensor: torch.Tensor, + spans: list[tuple[int, int]], + *, + head_group_num: int, + heads_per_group: int, + to_group_major: bool, +) -> None: + for start, seq_len in spans: + end = start + head_group_num * seq_len + leading_shape = ( + (seq_len, head_group_num) if to_group_major else (head_group_num, seq_len) + ) + tensor[start:end] = ( + tensor[start:end] + .clone() + .view(*leading_shape, heads_per_group, tensor.shape[-1]) + .transpose(0, 1) + .reshape(-1, heads_per_group, tensor.shape[-1]) + ) + + +def _copy_dense_page_table( + destination: torch.Tensor, + destination_row: int, + source: torch.Tensor, + source_row: int, + kv_len: int, + head_group_num: int, +) -> None: + for group in range(head_group_num): + destination[destination_row + group, :kv_len] = ( + source[source_row, :kv_len] * head_group_num + group + ) + + +def _gather_compressed_keys( + full_compressed_k: torch.Tensor, + level: CompressionLevelMetadata, + batches: list[int], +) -> tuple[torch.Tensor, torch.Tensor]: + lengths = [ + level.cu_seqlens_cpu[batch + 1] - level.cu_seqlens_cpu[batch] + for batch in batches + ] + compact_k = torch.cat( + [ + full_compressed_k[ + level.cu_seqlens_cpu[batch] : level.cu_seqlens_cpu[batch + 1] + ] + for batch in batches + ] + ) + compact_cu_seqlens = torch.tensor( + [0, *lengths], dtype=torch.int32, device=full_compressed_k.device + ).cumsum(0, dtype=torch.int32) + return compact_k, compact_cu_seqlens + + +def _copy_dense_page_tables( + metadata: MiniCPMSparseMetadata, + seq_lens_cpu, + head_group_num: int, +) -> None: + for dense_bs, row_start, _, _ in metadata.dense_layout: + _copy_dense_page_table( + metadata.sparse_page_table, + row_start, + metadata.base.page_table, + dense_bs, + int(seq_lens_cpu[dense_bs]), + head_group_num, + ) + + +class MiniCPMSparseBackend(AttentionBackend): + """MiniCPM sparse dispatch layered on the standard FlashAttention backend.""" + + def __init__( + self, + model_runner: ModelRunner, + skip_prefill: bool = False, + fa_impl_ver=3, + *, + use_flashinfer: bool, + ): + super().__init__() + use_blackwell = is_blackwell_supported() + if use_blackwell: + fa_impl_ver = 4 + self.flash_attn_backend = FlashAttentionBackend( + model_runner, + skip_prefill=skip_prefill, + fa_impl_ver=fa_impl_ver, + ) + self.forward_metadata: Optional[MiniCPMSparseMetadata] = None + self.max_context_len = self.flash_attn_backend.max_context_len + self.device = self.flash_attn_backend.device + self.model_dtype = model_runner.dtype + self._use_cuda_graph_buffers = False + self.decode_cuda_graph_metadata = ( + self.flash_attn_backend.decode_cuda_graph_metadata + ) + self.req_to_token_pool = self.flash_attn_backend.req_to_token_pool + self.token_to_kv_pool = self.flash_attn_backend.token_to_kv_pool + self.page_size = self.flash_attn_backend.page_size + tp_size = get_parallel().attn_tp_size + self.num_kv_heads = model_runner.model_config.get_num_kv_heads(tp_size) + self.num_q_heads = model_runner.model_config.num_attention_heads // tp_size + + # Sparse attention configuration (required for MiniCPM) + hf_config = model_runner.model_config.hf_config + + # MiniCPM must have sparse attention enabled + if not isinstance(hf_config, MiniCPMHybridConfig) or not ( + hf_config.has_minicpm_sparse_attention + ): + raise ValueError( + "MiniCPM model must have sparse attention enabled. " + "Please ensure the model config has MiniCPM sparse attention enabled." + ) + sparse_config = hf_config.sparse_config + self.kernel_size = sparse_config["kernel_size"] + self.kernel_stride = sparse_config["kernel_stride"] + self.init_blocks = sparse_config["init_blocks"] + self.block_size = sparse_config["block_size"] + self.window_size = sparse_config["window_size"] + if ( + self.kernel_stride <= 0 + or self.kernel_size <= 0 + or self.block_size <= 0 + or self.window_size < 0 + or self.kernel_size % self.kernel_stride + or self.block_size % self.kernel_stride + or self.window_size % self.block_size + ): + raise ValueError( + "MiniCPM sparse kernel_stride must divide kernel_size and " + "block_size, and block_size must divide window_size." + ) + attach_compressed_cache( + self.req_to_token_pool, + model_runner.token_to_kv_pool_allocator, + kernel_size=self.kernel_size, + kernel_stride=self.kernel_stride, + enable_memory_saver=model_runner.server_args.enable_memory_saver, + ) + self.req_to_sparse_k1_token = self.req_to_token_pool.req_to_sparse_k1_token + self.req_to_sparse_k2_token = self.req_to_token_pool.req_to_sparse_k2_token + self.minicpm_dense_as_sparse = envs.SGLANG_MINICPM_DENSE_AS_SPARSE.get() + self.dense_len = ( + 0 if self.minicpm_dense_as_sparse else sparse_config["dense_len"] + ) + self.config_dense_len = sparse_config["dense_len"] + topk = sparse_config["topk"] + self.local_blocks = self.window_size // self.block_size # local_blocks + self.sparse_topk = topk + (self.window_size // self.block_size) + self.num_sparse_topk_tokens = self.block_size * self.sparse_topk + required_context_len = max(self.config_dense_len, self.num_sparse_topk_tokens) + if self.max_context_len < required_context_len: + raise ValueError( + "MiniCPM sparse attention requires context_length >= " + f"{required_context_len}, got {self.max_context_len}." + ) + + # Head group number derived from model configuration + self.head_dim = model_runner.model_config.head_dim + self.head_group_num = self.num_kv_heads + self.heads_per_group = self.num_q_heads // self.head_group_num + if self.heads_per_group != 16: + raise ValueError( + "MiniCPM sparse attention requires 16 query heads per KV head, " + f"got {self.heads_per_group}." + ) + self.k1_kernel_size = self.kernel_size + self.k1_kernel_stride = self.kernel_stride + self.k2_kernel_size = self.kernel_size * 4 + self.k2_kernel_stride = self.kernel_stride * 4 + + self.minicpm_fuse_topk = ( + use_blackwell and use_flashinfer + ) or envs.SGLANG_MINICPM_FUSE_TOPK.get() + dtype_str = str(self.model_dtype).removeprefix("torch.") + if self.minicpm_fuse_topk and dtype_str not in ("bfloat16", "float16"): + raise ValueError( + "MiniCPM fused top-k only supports bfloat16 and float16, " + f"got {self.model_dtype}." + ) + + max_cache_len = self.max_context_len + pooled_k_len = (max_cache_len + self.block_size - 1) // self.block_size + + output_topk = min(self.sparse_topk, pooled_k_len) + + # For the kernel, we need power of 2 topk + topk_power2 = next_power_of_2(output_topk) + kernel_topk = min(topk_power2, pooled_k_len) + # Make sure it's still power of 2 + if kernel_topk != next_power_of_2(kernel_topk): + kernel_topk = next_power_of_2(kernel_topk) // 2 + kernel_topk = max(8, kernel_topk) + self.kernel_topk = kernel_topk + self.decode_fused_kernels = {} + self.prefill_fused_kernels = {} + bucketed_pooled_k_len = next_power_of_2(pooled_k_len) + + pooling_block_stride = self.block_size // self.kernel_stride # = 64 // 16 = 4 + pooling_pad_len = ( + self.kernel_size // self.kernel_stride - 1 + ) # = 32 // 16 - 1 = 1 + pooling_num_offs = ( + self.kernel_size // self.kernel_stride + + self.block_size // self.kernel_stride + - 1 + ) + self.fused_kernel_kwargs = { + "groups": self.heads_per_group, + "heads": self.num_q_heads, + "dim": self.head_dim, + "topk": self.kernel_topk, + "pooled_k_len": bucketed_pooled_k_len, + "m_block_dim": self.heads_per_group, + "block_M": self.heads_per_group, + "block_stride": pooling_block_stride, + "pad_len": pooling_pad_len, + "num_offs": pooling_num_offs, + "kernel_stride": self.kernel_stride, + "block_size": self.block_size, + "dense_len": self.dense_len, + "init_blocks": self.init_blocks, + "local_blocks": self.local_blocks, + "dtype_str": dtype_str, + } + chunked_prefill_size = model_runner.server_args.chunked_prefill_size + if self.minicpm_fuse_topk and chunked_prefill_size <= 0: + raise ValueError( + "MiniCPM fused top-k requires a positive --chunked-prefill-size." + ) + self.prefill_kernel_max_seqlen_q_grid = chunked_prefill_size + if self.minicpm_fuse_topk: + for batch_size in range(1, model_runner.max_running_requests + 1): + self._get_fused_topk_kernel(batch_size, is_prefill=True) + + self.attention_adapter = ( + MiniCPMFlashInferAdapter( + model_runner, + head_group_num=self.head_group_num, + heads_per_group=self.heads_per_group, + head_dim=self.head_dim, + page_size=self.page_size, + max_kv_tokens_per_row=max( + self.dense_len, + self.num_sparse_topk_tokens, + ), + ) + if use_flashinfer + else MiniCPMFlashAttentionAdapter(self.flash_attn_backend) + ) + + def _get_fused_topk_kernel(self, batch_size: int, *, is_prefill: bool): + if not self.minicpm_fuse_topk: + return None + + from sglang.srt.layers.attention.minicpm.fuse_kernel import ( + fused_attn_pooling_online_topk_decode, + fused_attn_pooling_online_topk_prefill, + ) + + cache = self.prefill_fused_kernels if is_prefill else self.decode_fused_kernels + if batch_size not in cache: + kwargs = dict(self.fused_kernel_kwargs, batch_size=batch_size) + if is_prefill: + kwargs["max_seqlen_q_grid"] = self.prefill_kernel_max_seqlen_q_grid + cache[batch_size] = fused_attn_pooling_online_topk_prefill(**kwargs) + else: + cache[batch_size] = fused_attn_pooling_online_topk_decode(**kwargs) + return cache[batch_size] + + def update_batch_for_sparse( + self, forward_batch: ForwardBatch, metadata: MiniCPMSparseMetadata + ): + cu_seqlens_q = metadata.base.cu_seqlens_q + + metadata.k1, metadata.k2 = _build_k1_k2_compression_metadata( + forward_batch=forward_batch, + base_metadata=metadata.base, + req_to_sparse_k1_token=self.req_to_sparse_k1_token, + req_to_sparse_k2_token=self.req_to_sparse_k2_token, + k1_kernel_size=self.k1_kernel_size, + k1_kernel_stride=self.k1_kernel_stride, + k2_kernel_size=self.k2_kernel_size, + k2_kernel_stride=self.k2_kernel_stride, + cu_seqlens_q=cu_seqlens_q, + ) + + if forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed(): + _plan_sparse_prefill( + forward_batch, + metadata, + head_group_num=self.head_group_num, + heads_per_group=self.heads_per_group, + dense_len=self.dense_len, + sparse_topk=self.sparse_topk, + block_size=self.block_size, + ) + _copy_dense_page_tables( + metadata, forward_batch.seq_lens_cpu, self.head_group_num + ) + else: + _plan_sparse_decode( + forward_batch=forward_batch, + metadata=metadata, + head_group_num=self.head_group_num, + dense_len=self.dense_len, + sparse_topk=self.sparse_topk, + block_size=self.block_size, + ) + + _copy_dense_page_tables( + metadata, forward_batch.seq_lens_cpu, self.head_group_num + ) + + # Stage1 optimization metadata for decode mode + metadata.cache_seqlens_int32_stage1 = ( + metadata.base.cache_seqlens_int32[metadata.sparse_bs_list] - 1 + ) + metadata.cu_seqlens_q_adjusted = ( + metadata.topk_cu_seqlens_q * self.heads_per_group + ) + metadata.max_seqlen_q_adjusted = self.heads_per_group + + def init_forward_metadata(self, forward_batch: ForwardBatch): + if forward_batch.forward_mode.is_target_verify(): + raise NotImplementedError( + "MiniCPM backend does not support speculative decoding (target verify)" + ) + if forward_batch.forward_mode.is_draft_extend_v2(): + raise NotImplementedError( + "MiniCPM backend does not support speculative decoding (draft extend)" + ) + + self._use_cuda_graph_buffers = False + self.flash_attn_backend.init_forward_metadata(forward_batch) + metadata = MiniCPMSparseMetadata(base=self.flash_attn_backend.forward_metadata) + if forward_batch.forward_mode.is_idle(): + self.forward_metadata = metadata + return + self.update_batch_for_sparse(forward_batch, metadata) + self.attention_adapter.prepare_forward( + metadata, + is_prefill=not forward_batch.forward_mode.is_decode_or_idle(), + graph=False, + ) + self.forward_metadata = metadata + + def _compress_decode_keys( + self, + query_states: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + ) -> tuple[torch.Tensor, torch.Tensor]: + metadata = self.forward_metadata + compressed = [] + for name, level in (("k1", metadata.k1), ("k2", metadata.k2)): + total = level.cu_seqlens_cpu[-1] + if self._use_cuda_graph_buffers: + buffer = self.decode_cuda_graph_metadata[f"compress_{name}"][:total] + else: + buffer = torch.empty( + (total, layer.tp_k_head_num, layer.head_dim), + dtype=query_states.dtype, + device=self.device, + ) + compressed.append(buffer) + compressed_k, compressed_k2 = compressed + + get_compress_k_v2( + layer=layer, + forward_batch=forward_batch, + metadata=metadata, + full_compressed_k1=compressed_k, + full_compressed_k2=compressed_k2, + max_context_length=self.max_context_len, + k1_kernel_size=self.k1_kernel_size, + k1_kernel_stride=self.k1_kernel_stride, + k2_kernel_size=self.k2_kernel_size, + k2_kernel_stride=self.k2_kernel_stride, + ) + return compressed_k, compressed_k2 + + def get_topk_for_sparse( + self, + query_states, + key_states, + layer, + forward_batch, + is_prefill=True, + ): + if is_prefill: + metadata = self.forward_metadata + sparse_bs = metadata.sparse_bs_list + full_compressed_k1, full_compressed_k2 = allocate_and_compress_keys( + layer=layer, + forward_batch=forward_batch, + metadata=metadata, + k1_token_nums=metadata.k1.cu_seqlens_cpu[-1], + k2_token_nums=metadata.k2.cu_seqlens_cpu[-1], + k1_kernel_size=self.k1_kernel_size, + k1_kernel_stride=self.k1_kernel_stride, + k2_kernel_size=self.k2_kernel_size, + k2_kernel_stride=self.k2_kernel_stride, + dtype=key_states.dtype, + device=key_states.device, + max_context_length=self.max_context_len, + ) + + compressed = [] + if len(sparse_bs) == forward_batch.batch_size: + compressed = [ + (full_compressed_k1, metadata.k1.cu_seqlens), + (full_compressed_k2, metadata.k2.cu_seqlens), + ] + else: + query_states = batched_gather( + query_states.reshape(-1, layer.tp_q_head_num, layer.head_dim), + forward_batch.extend_seq_lens_cpu, + sparse_bs, + ) + for full_compressed_k, level in ( + (full_compressed_k1, metadata.k1), + (full_compressed_k2, metadata.k2), + ): + compressed.append( + _gather_compressed_keys(full_compressed_k, level, sparse_bs) + ) + + (compressed_k, compressed_cu_seqlens), ( + compressed_k2, + compressed_cu_seqlens2, + ) = compressed + + ret = self.sparse_get_topk_impl( + query_states, + metadata.topk_cu_seqlens_q, + metadata.topk_cu_seqlens_k, + metadata.topk_max_seqlen_q, + metadata.topk_max_seqlen_k, + compressed_k=compressed_k, + compressed_cu_seqlens=compressed_cu_seqlens, + compressed_k2=compressed_k2, + compressed_cu_seqlens2=compressed_cu_seqlens2, + fused_kernel=self._get_fused_topk_kernel( + len(sparse_bs), + is_prefill=True, + ), + ) + return ret + else: + metadata = self.forward_metadata + compressed_k, compressed_k2 = self._compress_decode_keys( + query_states, + layer, + forward_batch, + ) + + sparse_bs = metadata.sparse_bs_list + if not sparse_bs: + return None + + cu_seqlens_q = metadata.base.cu_seqlens_q + compressed_cu_seqlens = metadata.k1.cu_seqlens + compressed_cu_seqlens2 = metadata.k2.cu_seqlens + if len(sparse_bs) < forward_batch.batch_size: + query_states = query_states[sparse_bs] + compressed_k, compressed_cu_seqlens = _gather_compressed_keys( + compressed_k, metadata.k1, sparse_bs + ) + compressed_k2, compressed_cu_seqlens2 = _gather_compressed_keys( + compressed_k2, metadata.k2, sparse_bs + ) + cu_seqlens_q = metadata.topk_cu_seqlens_q + + ret = self.sparse_get_topk_impl( + query_states, + cu_seqlens_q, + metadata.base.cu_seqlens_k, + 1, + metadata.base.max_seq_len_k, + compressed_k=compressed_k, + compressed_cu_seqlens=compressed_cu_seqlens, + compressed_k2=compressed_k2, + compressed_cu_seqlens2=compressed_cu_seqlens2, + fused_kernel=self._get_fused_topk_kernel( + len(sparse_bs), + is_prefill=False, + ), + ) + + return ret + + def sparse_get_topk_impl( + self, + query_layer, + cu_seqlens_q, + cu_seqlens_k, + max_seqlen_in_batch_q, + max_seqlen_in_batch_k, + compressed_k=None, + compressed_cu_seqlens=None, + compressed_k2=None, + compressed_cu_seqlens2=None, + fused_kernel=None, + ): + cache_lens = None + if max_seqlen_in_batch_k > max_seqlen_in_batch_q: + if max_seqlen_in_batch_q == 1: + cache_lens = self.forward_metadata.cache_seqlens_int32_stage1 + else: + seq_lens_k = cu_seqlens_k[1:] - cu_seqlens_k[:-1] + seq_lens_q = cu_seqlens_q[1:] - cu_seqlens_q[:-1] + cache_lens = seq_lens_k - seq_lens_q + else: + batch_size = cu_seqlens_q.shape[0] - 1 + cache_lens = torch.zeros( + batch_size, dtype=torch.int32, device=cu_seqlens_q.device + ) + + if not self.minicpm_fuse_topk: + topk_idx = compressed_attention( + query_layer, + compressed_k, + compressed_k2, + self.kernel_stride, + self.block_size, + self.sparse_topk, + cu_seqlens_q, + compressed_cu_seqlens, + compressed_cu_seqlens2, + max_seqlen_in_batch_q, + self.max_context_len, + init_blocks=self.init_blocks, + local_blocks=self.local_blocks, + cache_lens=cache_lens, + cu_seqlens_q_adjusted=self.forward_metadata.cu_seqlens_q_adjusted, + max_seqlen_q_adjusted=self.forward_metadata.max_seqlen_q_adjusted, + ) + else: + topk_idx = compressed_attention_tilelang( + query_layer, + compressed_k, + self.block_size, + self.sparse_topk, + self.kernel_topk, + cu_seqlens_q, + compressed_cu_seqlens, + cache_lens=cache_lens, + fused_kernel=fused_kernel, + max_cache_len=self.max_context_len, + ) + + return topk_idx + + def forward_extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache=True, + # For multi-head latent attention + q_rope: Optional[torch.Tensor] = None, + k_rope: Optional[torch.Tensor] = None, + sinks: Optional[torch.Tensor] = None, + ): + if layer.is_cross_attention: + raise NotImplementedError( + "MiniCPM backend does not support cross attention" + ) + if layer.sliding_window_size not in (None, -1): + raise NotImplementedError( + "MiniCPM backend does not support sliding-window attention" + ) + if forward_batch.forward_mode.is_draft_extend_v2(): + raise NotImplementedError( + "MiniCPM backend does not support draft extend mode" + ) + + if k is not None: + assert v is not None + if save_kv_cache: + self.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k, + v, + layer.k_scale, + layer.v_scale, + ) + + metadata = self.forward_metadata + q, q_rope, k_rope, k_descale, v_descale = ( + self.flash_attn_backend.prepare_paged_mha_query( + q, + q_rope, + k_rope, + layer, + logical_batch_size=forward_batch.batch_size, + kv_head_num=layer.tp_k_head_num, + is_prefill=True, + ) + ) + page_table = metadata.base.page_table + + if metadata.sparse_bs_list: + q_reshaped = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim) + topk_idx = self.get_topk_for_sparse( + query_states=q_reshaped, + key_states=k, + layer=layer, + forward_batch=forward_batch, + ) + + sparse_page_table_sparse_bs = get_block_table( + topk_idx, + page_table[metadata.sparse_bs_list], + metadata.token_to_bs, + metadata.token_pos_in_bs, + metadata.seqlen_k_sparse_bs_tensor, + head_group_num=self.head_group_num, + block_size=self.block_size, + elementwise=False, + ).reshape(-1, self.num_sparse_topk_tokens) + + # copy page table for sparse bs + metadata.sparse_page_table[ + metadata.sparse_idx, : self.num_sparse_topk_tokens + ] = sparse_page_table_sparse_bs + else: + total_k1 = self.forward_metadata.k1.cu_seqlens_cpu[-1] + total_k2 = self.forward_metadata.k2.cu_seqlens_cpu[-1] + + allocate_and_compress_keys( + layer=layer, + forward_batch=forward_batch, + metadata=self.forward_metadata, + k1_token_nums=total_k1, + k2_token_nums=total_k2, + k1_kernel_size=self.k1_kernel_size, + k1_kernel_stride=self.k1_kernel_stride, + k2_kernel_size=self.k2_kernel_size, + k2_kernel_stride=self.k2_kernel_stride, + dtype=k.dtype, + device=k.device, + max_context_length=self.max_context_len, + ) + + dense_layout_spans = [ + (query_start, query_len) + for _, _, query_start, query_len in metadata.dense_layout + ] + + q_by_head_group = q.contiguous().view(-1, self.heads_per_group, layer.head_dim) + _transpose_head_group_layout( + q_by_head_group, + dense_layout_spans, + head_group_num=self.head_group_num, + heads_per_group=self.heads_per_group, + to_group_major=True, + ) + + key_cache, value_cache = self.flash_attn_backend.get_paged_mha_kv_cache( + layer, + head_group_num=self.head_group_num, + ) + + result = self.attention_adapter.forward( + q_by_head_group, + key_cache, + value_cache, + metadata, + layer, + is_prefill=True, + k_descale=k_descale, + v_descale=v_descale, + sinks=sinks, + ) + + _transpose_head_group_layout( + result, + dense_layout_spans, + head_group_num=self.head_group_num, + heads_per_group=self.heads_per_group, + to_group_major=False, + ) + + return result.view(-1, layer.tp_q_head_num * layer.head_dim) + + def forward_decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache=True, + # For multi-head latent attention + q_rope: Optional[torch.Tensor] = None, + k_rope: Optional[torch.Tensor] = None, + sinks: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if layer.is_cross_attention: + raise NotImplementedError( + "MiniCPM backend does not support cross attention" + ) + if layer.sliding_window_size not in (None, -1): + raise NotImplementedError( + "MiniCPM backend does not support sliding-window attention" + ) + + bs = forward_batch.batch_size + if k is not None: + assert v is not None + if save_kv_cache: + self.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k, + v, + layer.k_scale, + layer.v_scale, + ) + + metadata = self.forward_metadata + q, q_rope, k_rope, k_descale, v_descale = ( + self.flash_attn_backend.prepare_paged_mha_query( + q, + q_rope, + k_rope, + layer, + logical_batch_size=bs, + kv_head_num=layer.tp_k_head_num, + is_prefill=False, + ) + ) + key_cache, value_cache = self.flash_attn_backend.get_paged_mha_kv_cache( + layer, + head_group_num=self.head_group_num, + ) + + page_table = metadata.base.page_table + cache_seqlens = metadata.base.cache_seqlens_int32 + q_reshaped = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim) + + topk_idx = self.get_topk_for_sparse( + query_states=q_reshaped, + key_states=k, + layer=layer, + forward_batch=forward_batch, + is_prefill=False, + ) + if topk_idx is not None: + topk_page_table = page_table + topk_cache_seqlens = cache_seqlens + if not self._use_cuda_graph_buffers: + topk_page_table = page_table[metadata.sparse_bs_list] + topk_cache_seqlens = cache_seqlens[metadata.sparse_bs_list] + sparse_page_table = get_block_table( + topk_idx, + topk_page_table, + metadata.token_to_bs, + topk_cache_seqlens, + topk_cache_seqlens, + head_group_num=self.head_group_num, + block_size=self.block_size, + elementwise=True, + ).reshape(-1, self.num_sparse_topk_tokens) + destination = metadata.sparse_page_table[:, : self.num_sparse_topk_tokens] + if self._use_cuda_graph_buffers: + destination.copy_( + torch.where( + (cache_seqlens >= self.dense_len).repeat_interleave( + self.head_group_num + )[:, None], + sparse_page_table, + destination, + ) + ) + else: + destination[metadata.sparse_idx] = sparse_page_table + + q_reshaped_by_head_group = q_reshaped.reshape( + -1, self.heads_per_group, layer.head_dim + ) + assert self.page_size == 1 + result = self.attention_adapter.forward( + q_reshaped_by_head_group, + key_cache, + value_cache, + metadata, + layer, + is_prefill=False, + k_descale=k_descale, + v_descale=v_descale, + sinks=sinks, + ) + + return result.view(-1, layer.tp_q_head_num * layer.v_head_dim) + + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): + self.flash_attn_backend.init_cuda_graph_state(max_bs, max_num_tokens) + self.attention_adapter.init_cuda_graph_state(max_bs * self.head_group_num) + buffers = self.flash_attn_backend.decode_cuda_graph_metadata + self.decode_cuda_graph_metadata = buffers + sparse_max_num_pages = ( + max(self.dense_len, self.num_sparse_topk_tokens) + self.page_size - 1 + ) // self.page_size + buffers.update( + { + "sparse_cache_seqlens": torch.full( + (max_bs * self.head_group_num,), + self.num_sparse_topk_tokens, + dtype=torch.int32, + device=self.device, + ), + "sparse_cu_seqlens_q": torch.arange( + 0, + max_bs * self.head_group_num + 1, + dtype=torch.int32, + device=self.device, + ), + "sparse_cu_seqlens_k": torch.arange( + 0, + (max_bs * self.head_group_num + 1) * self.num_sparse_topk_tokens, + self.num_sparse_topk_tokens, + dtype=torch.int32, + device=self.device, + ), + "token_to_bs": torch.arange( + 0, max_bs, dtype=torch.int32, device=self.device + ), + "sparse_page_table": torch.zeros( + max_bs * self.head_group_num, + sparse_max_num_pages, + dtype=torch.int32, + device=self.device, + ), + "cu_seqlens_q_adjusted": torch.arange( + 0, max_bs + 1, dtype=torch.int32, device=self.device + ) + * self.heads_per_group, + "cache_seqlens_int32_stage1": torch.zeros( + max_bs, dtype=torch.int32, device=self.device + ), + } + ) + + for name, kernel_size, kernel_stride in ( + ("k1", self.k1_kernel_size, self.k1_kernel_stride), + ("k2", self.k2_kernel_size, self.k2_kernel_stride), + ): + max_num_pages = ( + max( + 0, + (self.max_context_len - kernel_size) // kernel_stride + 1, + ) + + self.page_size + - 1 + ) // self.page_size + buffers[f"compress_{name}"] = torch.zeros( + ( + max_bs * self.max_context_len // kernel_stride, + self.head_group_num, + self.head_dim, + ), + dtype=self.model_dtype, + device=self.device, + ) + buffers[f"{name}.table"] = torch.zeros( + max_bs, max_num_pages, dtype=torch.int32, device=self.device + ) + buffers[f"{name}.history_compress_token_nums"] = torch.zeros( + max_bs, dtype=torch.int32, device=self.device + ) + for field in ( + "cu_seqlens", + "cu_new_token_nums", + "cu_total_compress_token_nums", + ): + buffers[f"{name}.{field}"] = torch.zeros( + max_bs + 1, dtype=torch.int32, device=self.device + ) + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + if not forward_batch.forward_mode.is_decode_or_idle(): + raise NotImplementedError( + "MiniCPM backend CUDA graph only supports decode/idle mode, " + f"got {forward_batch.forward_mode}" + ) + + self._use_cuda_graph_buffers = True + self._get_fused_topk_kernel( + forward_batch.batch_size, + is_prefill=False, + ) + self.flash_attn_backend.init_forward_metadata_out_graph( + forward_batch, in_capture + ) + metadata = MiniCPMSparseMetadata(base=self.flash_attn_backend.forward_metadata) + self._bind_sparse_graph_metadata( + forward_batch, + metadata, + in_capture=in_capture, + ) + if not in_capture: + self._replay_sparse_graph_metadata(forward_batch, metadata) + self.attention_adapter.prepare_forward( + metadata, + is_prefill=False, + graph=True, + ) + self.forward_metadata = metadata + + def _build_sparse_decode_replay_metadata( + self, + forward_batch: ForwardBatch, + metadata: MiniCPMSparseMetadata, + ): + decode_metadata = MiniCPMSparseMetadata(base=metadata.base) + _plan_sparse_decode( + forward_batch=forward_batch, + metadata=decode_metadata, + head_group_num=self.head_group_num, + dense_len=self.dense_len, + sparse_topk=self.sparse_topk, + block_size=self.block_size, + ) + compression_metadata = _build_k1_k2_compression_metadata( + forward_batch=forward_batch, + base_metadata=metadata.base, + req_to_sparse_k1_token=self.req_to_sparse_k1_token, + req_to_sparse_k2_token=self.req_to_sparse_k2_token, + k1_kernel_size=self.k1_kernel_size, + k1_kernel_stride=self.k1_kernel_stride, + k2_kernel_size=self.k2_kernel_size, + k2_kernel_stride=self.k2_kernel_stride, + cu_seqlens_q=metadata.base.cu_seqlens_q, + ) + return decode_metadata, compression_metadata + + def _bind_sparse_graph_metadata( + self, + forward_batch: ForwardBatch, + metadata: MiniCPMSparseMetadata, + *, + in_capture: bool, + ): + bs = forward_batch.batch_size + buffers = self.decode_cuda_graph_metadata + sparse_rows = self.head_group_num * bs + metadata.sparse_cache_seqlens_int32 = buffers["sparse_cache_seqlens"][ + :sparse_rows + ] + metadata.sparse_cu_seqlens_q = buffers["sparse_cu_seqlens_q"][: sparse_rows + 1] + metadata.sparse_cu_seqlens_k = buffers["sparse_cu_seqlens_k"][: sparse_rows + 1] + metadata.token_to_bs = buffers["token_to_bs"][:bs] + metadata.sparse_page_table = buffers["sparse_page_table"][:sparse_rows] + metadata.sparse_bs_list = list(range(bs)) + + assume_kv_len = self.config_dense_len + if in_capture: + metadata.base.cu_seqlens_k.copy_( + torch.arange(bs + 1, device=self.device, dtype=torch.int32) + * assume_kv_len + ) + metadata.base.max_seq_len_k = assume_kv_len + + for name, kernel_size, kernel_stride in ( + ("k1", self.k1_kernel_size, self.k1_kernel_stride), + ("k2", self.k2_kernel_size, self.k2_kernel_stride), + ): + level = CompressionLevelMetadata() + setattr(metadata, name, level) + level_len = max(0, (assume_kv_len - kernel_size) // kernel_stride + 1) + level.cu_seqlens_cpu = [index * level_len for index in range(bs + 1)] + level.cu_seqlens = buffers[f"{name}.cu_seqlens"][: bs + 1] + if in_capture: + level.cu_seqlens.copy_( + torch.arange(bs + 1, device=self.device, dtype=torch.int32) + * level_len + ) + level.table = buffers[f"{name}.table"][:bs] + level.history_compress_token_nums = buffers[ + f"{name}.history_compress_token_nums" + ][:bs] + for field in ( + "cu_new_token_nums", + "cu_total_compress_token_nums", + ): + setattr(level, field, buffers[f"{name}.{field}"][: bs + 1]) + + metadata.cu_seqlens_q_adjusted = buffers["cu_seqlens_q_adjusted"][: bs + 1] + metadata.cache_seqlens_int32_stage1 = buffers["cache_seqlens_int32_stage1"][:bs] + metadata.max_seqlen_q_adjusted = ( + metadata.base.max_seq_len_q * self.heads_per_group + ) + + def _replay_sparse_graph_metadata( + self, + forward_batch: ForwardBatch, + metadata: MiniCPMSparseMetadata, + ): + bs = forward_batch.batch_size + real_bs = bs - forward_batch.num_padding + if real_bs == 0: + metadata.sparse_cache_seqlens_int32.zero_() + metadata.sparse_cu_seqlens_k.zero_() + metadata.cache_seqlens_int32_stage1.zero_() + for level in (metadata.k1, metadata.k2): + level.history_compress_token_nums.zero_() + level.cu_seqlens.zero_() + level.cu_new_token_nums.zero_() + level.cu_total_compress_token_nums.zero_() + return + + sparse_forward_batch = SimpleNamespace( + batch_size=real_bs, + req_pool_indices=forward_batch.req_pool_indices[:real_bs], + seq_lens_cpu=forward_batch.seq_lens_cpu[:real_bs], + ) + decode_metadata, compression_metadata = ( + self._build_sparse_decode_replay_metadata(sparse_forward_batch, metadata) + ) + real_sparse_rows = self.head_group_num * real_bs + metadata.sparse_cache_seqlens_int32[:real_sparse_rows].copy_( + decode_metadata.sparse_cache_seqlens_int32 + ) + metadata.sparse_cu_seqlens_k[: real_sparse_rows + 1].copy_( + decode_metadata.sparse_cu_seqlens_k + ) + metadata.cache_seqlens_int32_stage1[:real_bs].copy_( + metadata.base.cache_seqlens_int32[:real_bs] - 1 + ) + metadata.dense_layout = decode_metadata.dense_layout + _copy_dense_page_tables( + metadata, forward_batch.seq_lens_cpu, self.head_group_num + ) + + for (name, req_to_sparse), src in zip( + ( + ("k1", self.req_to_sparse_k1_token), + ("k2", self.req_to_sparse_k2_token), + ), + compression_metadata, + ): + dst = getattr(metadata, name) + dst.history_compress_token_nums[:real_bs].copy_( + src.history_compress_token_nums + ) + if real_bs < bs: + dst.history_compress_token_nums[real_bs:].zero_() + for field in ( + "cu_seqlens", + "cu_new_token_nums", + "cu_total_compress_token_nums", + ): + dst_field = getattr(dst, field) + src_field = getattr(src, field) + dst_field[: real_bs + 1].copy_(src_field) + if real_bs < bs: + dst_field[real_bs + 1 :].fill_(src_field[-1]) + dst.table.copy_(req_to_sparse[forward_batch.req_pool_indices]) + + if real_bs < bs: + metadata.sparse_cache_seqlens_int32[real_sparse_rows:].zero_() + metadata.sparse_cu_seqlens_k[real_sparse_rows + 1 :].fill_( + decode_metadata.sparse_cu_seqlens_k[-1] + ) + metadata.cache_seqlens_int32_stage1[real_bs:].zero_() + + def get_cuda_graph_seq_len_fill_value(self): + return self.flash_attn_backend.get_cuda_graph_seq_len_fill_value() diff --git a/python/sglang/srt/layers/attention/minicpm/cache.py b/python/sglang/srt/layers/attention/minicpm/cache.py new file mode 100644 index 000000000..971bcc523 --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/cache.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE +from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter + +if TYPE_CHECKING: + from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + + +class MiniCPMCompressedCache: + def __init__( + self, + pool: ReqToTokenPool, + allocator: BaseTokenToKVPoolAllocator, + *, + kernel_size: int, + kernel_stride: int, + enable_memory_saver: bool, + ): + self.pool = pool + self.allocator = allocator + self.kernel_size = kernel_size + self.kernel_stride = kernel_stride + saver = TorchMemorySaverAdapter.create(enable=enable_memory_saver) + with saver.region(GPU_MEMORY_TYPE_KV_CACHE): + k1_size = (pool.max_context_len - kernel_size) // kernel_stride + 1 + k2_size = (pool.max_context_len - kernel_size * 4) // ( + kernel_stride * 4 + ) + 1 + pool.req_to_sparse_k1_token = torch.zeros( + (pool._alloc_size, k1_size), dtype=torch.int32, device=pool.device + ) + pool.req_to_sparse_k2_token = torch.zeros( + (pool._alloc_size, k2_size), dtype=torch.int32, device=pool.device + ) + self.allocated_lens = [ + [0] * pool._alloc_size, + [0] * pool._alloc_size, + ] + self.reserved_slots = torch.empty(0, dtype=torch.int64, device=pool.device) + self.free_slots = self.reserved_slots + self.reset_allocator() + + def reset_allocator(self) -> None: + """Reserve K1/K2 capacity after the backing allocator is cleared.""" + if self.allocator.page_size != 1: + raise ValueError("MiniCPM sparse attention requires page_size=1") + + total_slots = self.allocator.available_size() + # K1 and K2 consume at most 1/s and 1/(4s) slots per dense token. + denominator = 4 * self.kernel_stride + 5 + reserve_size = (5 * total_slots + denominator - 1) // denominator + reserved_slots = self.allocator.alloc(reserve_size) + if reserved_slots is None: + raise RuntimeError( + f"Unable to reserve {reserve_size} MiniCPM compressed-cache slots" + ) + + self.dense_capacity = total_slots - reserve_size + self.reserved_slots = reserved_slots + self.free_slots = reserved_slots + self.clear() + + def _alloc_reserved(self, size: int) -> torch.Tensor: + if size > len(self.free_slots): + raise RuntimeError( + "MiniCPM compressed cache is out of reserved slots: " + f"requested={size}, available={len(self.free_slots)}" + ) + slots = self.free_slots[:size] + self.free_slots = self.free_slots[size:] + return slots + + def _free_reserved(self, slots: torch.Tensor) -> None: + self.free_slots = torch.cat( + (self.free_slots, slots.to(self.reserved_slots.dtype)) + ) + + def _sparse_len(self, length: int, scale: int) -> int: + kernel_size = self.kernel_size * scale + if length < kernel_size: + return 0 + return (length - kernel_size) // (self.kernel_stride * scale) + 1 + + def alloc_to_lengths( + self, + *, + req_pool_indices_cpu: torch.Tensor, + target_seq_lens_cpu: torch.Tensor, + ) -> None: + req_indices = req_pool_indices_cpu.tolist() + seq_lens = target_seq_lens_cpu.tolist() + tables = ( + self.pool.req_to_sparse_k1_token, + self.pool.req_to_sparse_k2_token, + ) + plans = [] + for level, (table, scale) in enumerate(zip(tables, (1, 4))): + targets = { + req_idx: self._sparse_len(seq_len, scale) + for req_idx, seq_len in zip(req_indices, seq_lens) + } + rows = [ + (req_idx, self.allocated_lens[level][req_idx], target) + for req_idx, target in targets.items() + if target > self.allocated_lens[level][req_idx] + ] + plans.append((table, rows, sum(end - start for _, start, end in rows))) + + allocated = [] + try: + for _, _, size in plans: + allocated.append(self._alloc_reserved(size) if size > 0 else None) + + for (table, rows, _), locs in zip(plans, allocated): + if locs is None: + continue + offset = 0 + for req_idx, start, end in rows: + count = end - start + table[req_idx, start:end] = locs[offset : offset + count].to( + torch.int32 + ) + offset += count + for level, (_, rows, _) in enumerate(plans): + for req_idx, _, end in rows: + self.allocated_lens[level][req_idx] = end + except Exception: + for locs in allocated: + if locs is not None: + self._free_reserved(locs) + raise + + def free(self, req_pool_idx: int) -> None: + allocated = [] + for table, lengths in zip( + ( + self.pool.req_to_sparse_k1_token, + self.pool.req_to_sparse_k2_token, + ), + self.allocated_lens, + ): + length = lengths[req_pool_idx] + if length > 0: + allocated.append(table[req_pool_idx, :length].clone()) + table[req_pool_idx, :length].zero_() + lengths[req_pool_idx] = 0 + + if allocated: + self._free_reserved(torch.cat(allocated)) + + def clear(self) -> None: + self.pool.req_to_sparse_k1_token.zero_() + self.pool.req_to_sparse_k2_token.zero_() + for lengths in self.allocated_lens: + lengths[:] = [0] * len(lengths) + self.free_slots = self.reserved_slots + + +def attach_compressed_cache( + pool: ReqToTokenPool, + allocator: BaseTokenToKVPoolAllocator, + *, + kernel_size: int, + kernel_stride: int, + enable_memory_saver: bool, +) -> ReqToTokenPool: + if isinstance(pool._aux_cache, MiniCPMCompressedCache): + return pool + + pool.attach_aux_cache( + MiniCPMCompressedCache( + pool, + allocator, + kernel_size=kernel_size, + kernel_stride=kernel_stride, + enable_memory_saver=enable_memory_saver, + ) + ) + return pool diff --git a/python/sglang/srt/layers/attention/minicpm/fuse_kernel.py b/python/sglang/srt/layers/attention/minicpm/fuse_kernel.py new file mode 100644 index 000000000..7a1864edd --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/fuse_kernel.py @@ -0,0 +1,425 @@ +import math +from functools import partial + +import tilelang +import tilelang.language as T +import tilelang.math + +_pass_configs = { + tilelang.PassConfigKey.TL_DISABLE_THREAD_STORAGE_SYNC: True, + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, +} + + +@tilelang.jit(pass_configs=_pass_configs) +def _fused_attn_pooling_online_topk( + batch_size: int, + groups: int, + heads: int, + dim: int, + topk: int, + max_seqlen_q_grid: int, # Static param for grid (use bucketing) + pooled_k_len: int, # Static param (use bucketing) = ceil(max_seqlen_k / block_size) + is_causal: bool, + dense_len: int = 0, + m_block_dim: int = 16, + block_M: int = 16, + block_N: int = 64, + # infllmv2 pooling parameters: + # block_stride = block_size // kernel_stride = 64 // 16 = 4 + # pad_len = kernel_size // kernel_stride - 1 = 32 // 16 - 1 = 1 + # num_offs = kernel_size // kernel_stride + block_size // kernel_stride - 1 = 2 + 4 - 1 = 5 + block_stride: int = 4, # pool output block stride + pad_len: int = 1, # padding for pool blocks + num_offs: int = 5, # number of k positions each pool block reads + kernel_stride: int = 16, + block_size: int = 64, # block size for q/k block computation + init_blocks: int = 0, + local_blocks: int = 0, + num_stages: int = 0, + threads: int = 128, + dtype_str: str = "bfloat16", +): + """ + Fused Attention + Max Pooling + Online TopK for prefill and decode. + + Chunk prefill support: + - cache_lens: tensor of shape [batch_size], cache length for each batch + - When cache_lens[i] = 0, it's standard prefill + - When cache_lens[i] > 0, it's chunk prefill (continuing from cached state) + + Pooling logic aligned with infllmv2_cuda_impl: + - For each pool block b, it aggregates k scores in range [b * block_stride - pad_len, b * block_stride - pad_len + num_offs) + - block_stride = 4, pad_len = 1, num_offs = 5 + - pool block 0: k in [0-1, 0-1+5) = [-1, 4) -> [0, 4) + - pool block 1: k in [3, 8) + - pool block 2: k in [7, 12) + - etc. + """ + assert topk == tilelang.math.next_power_of_2(topk), "topk must be power of 2" + + scale = (1.0 / dim) ** 0.5 * 1.44269504 + head_kv = heads // groups + + # Dynamic dimensions - inferred from tensor shapes at runtime + UQ = T.dynamic("UQ") + UKV = T.dynamic("UKV") + + q_shape = [UQ * groups, head_kv, dim] + kv_shape = [UKV, head_kv, dim] + topk_indices_shape = [head_kv, UQ, topk] + topk_values_shape = [head_kv, UQ, topk] + + dtype = dtype_str + accum_dtype = "float" + + N = 2 * topk + num_sort_iters = int(round(math.log2(N))) + block_P = topk + + @T.macro + def bitonic_sort( + topk_index_shared: T.Buffer([N], "int32"), + topk_value_shared: T.Buffer([N], "float32"), + ): + T.sync_threads() + for i1 in T.serial(num_sort_iters): + for i2 in T.serial(i1 + 1): + for i in T.Parallel(N): + ascending = (i & (1 << (i1 + 1))) != 0 + j = i ^ (1 << (i1 - i2)) + if i < j and ( + (ascending and topk_value_shared[i] > topk_value_shared[j]) + or ( + not ascending + and topk_value_shared[i] < topk_value_shared[j] + ) + ): + val = topk_value_shared[i] + topk_value_shared[i] = topk_value_shared[j] + topk_value_shared[j] = val + idx = topk_index_shared[i] + topk_index_shared[i] = topk_index_shared[j] + topk_index_shared[j] = idx + T.sync_threads() + + @T.prim_func + def main( + Q_unpad: T.Tensor(q_shape, dtype), + K_unpad: T.Tensor(kv_shape, dtype), + cu_seqlens_q: T.Tensor([batch_size + 1], "int32"), + cu_seqlens_k: T.Tensor([batch_size + 1], "int32"), + cache_lens: T.Tensor( + [batch_size], "int32" + ), # Per-batch cache length for chunk prefill + TopkIndices: T.Tensor(topk_indices_shape, "int32"), + TopkValues: T.Tensor(topk_values_shape, "float32"), + ): + with T.Kernel(max_seqlen_q_grid, head_kv, batch_size, threads=threads) as ( + bx, + by, + bz, + ): + Q_shared = T.alloc_shared([block_M, dim], dtype) + K_shared = T.alloc_shared([block_N, dim], dtype) + topk_index_shared = T.alloc_shared([N], "int32") + topk_value_shared = T.alloc_shared([N], "float32") + pool_max_shared = T.alloc_shared([block_P], "float32") + + acc_s = T.alloc_fragment([block_M, block_N], accum_dtype) + scores_max = T.alloc_fragment([block_M], accum_dtype) + scores_max_prev = T.alloc_fragment([block_M], accum_dtype) + scores_scale = T.alloc_fragment([block_M], accum_dtype) + scores_sum = T.alloc_fragment([block_M], accum_dtype) + logsum = T.alloc_fragment([block_M], accum_dtype) + acc_output = T.alloc_fragment([block_N], accum_dtype) + + batch_idx = bz + kv_head_idx = by + original_q_idx = bx + + q_start_idx = cu_seqlens_q[batch_idx] + k_start_idx = cu_seqlens_k[batch_idx] + q_end_idx = cu_seqlens_q[batch_idx + 1] + k_end_idx = cu_seqlens_k[batch_idx + 1] + + q_current_seqlen = T.alloc_var("int32", init=q_end_idx - q_start_idx) + k_current_seqlen = T.alloc_var("int32", init=k_end_idx - k_start_idx) + + # Chunk prefill: cache_len from tensor (0 for standard prefill, >0 for chunk prefill) + cache_len = cache_lens[batch_idx] + if not is_causal: + active = cache_len + 1 >= dense_len + q_current_seqlen = T.if_then_else(active, q_current_seqlen, 0) + k_current_seqlen = T.if_then_else(active, k_current_seqlen, 0) + if is_causal: + actual_pooled_k_len = ( + k_current_seqlen - 1 + pad_len + ) // block_stride + 1 + else: + actual_pooled_k_len = (1 + cache_len + block_size - 1) // block_size + effective_pooled_k_len = T.min(actual_pooled_k_len, pooled_k_len) + + T.fill(topk_index_shared, -1) + T.fill(topk_value_shared, float("-inf")) + T.sync_threads() + + # Use q_end_idx to avoid out-of-bounds access for Q + q_copy_end = T.min( + q_start_idx * groups + (bx + 1) * block_M, q_end_idx * groups + ) + T.copy( + Q_unpad[ + q_start_idx * groups + bx * block_M : q_copy_end, kv_head_idx, : + ], + Q_shared, + ) + for i, d in T.Parallel(block_M, dim): + if original_q_idx >= q_current_seqlen: + Q_shared[i, d] = 0 + + T.fill(logsum, 0) + T.fill(scores_max, -T.infinity(accum_dtype)) + + loop_range_k = T.ceildiv(k_current_seqlen, block_N) + + for k in T.Pipelined(loop_range_k, num_stages=num_stages): + # Use k_end_idx to avoid out-of-bounds access for last block + k_copy_end = T.min(k_start_idx + (k + 1) * block_N, k_end_idx) + T.copy( + K_unpad[k_start_idx + k * block_N : k_copy_end, kv_head_idx, :], + K_shared, + ) + for i, d in T.Parallel(block_N, dim): + if k * block_N + i >= k_current_seqlen: + K_shared[i, d] = 0 + + for i, j in T.Parallel(block_M, block_N): + k_idx = k * block_N + j + boundary_mask = (original_q_idx >= q_current_seqlen) or ( + k_idx >= k_current_seqlen + ) + if is_causal: + row_idx = original_q_idx * block_M + i + cache_len * block_M + orig_row_idx = row_idx // m_block_dim + orig_seqlen_q = ( + (q_current_seqlen + cache_len) * block_M + ) // m_block_dim + compressed_seqlen_q = ( + orig_seqlen_q - kernel_stride + 1 + ) // kernel_stride + offset_row_idx = T.max( + 0, + (orig_row_idx + 1) // kernel_stride + - 1 + + k_current_seqlen + - compressed_seqlen_q, + ) + q_compress_clamped = T.min(k_current_seqlen, offset_row_idx) + causal_mask = k_idx > q_compress_clamped + acc_s[i, j] = T.if_then_else( + boundary_mask or causal_mask, -1e9, 0 + ) + else: + acc_s[i, j] = T.if_then_else(boundary_mask, -1e9, 0) + + T.gemm( + Q_shared, + K_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + + T.copy(scores_max, scores_max_prev) + T.fill(scores_max, -T.infinity(accum_dtype)) + T.reduce_max(acc_s, scores_max, dim=1, clear=False) + for i in T.Parallel(block_M): + scores_max[i] = T.max(scores_max[i], scores_max_prev[i]) + + for i in T.Parallel(block_M): + scores_scale[i] = T.exp2( + scores_max_prev[i] * scale - scores_max[i] * scale + ) + + for i, j in T.Parallel(block_M, block_N): + acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale) + + T.reduce_sum(acc_s, scores_sum, dim=1) + for i in T.Parallel(block_M): + logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i] + + loop_range_pool = T.ceildiv(effective_pooled_k_len, block_P) + + for p_block in T.serial(loop_range_pool): + T.fill(pool_max_shared, float("-inf")) + T.sync_threads() + + for k in T.serial(loop_range_k): + # Use k_end_idx to avoid out-of-bounds access for last block + k_copy_end = T.min(k_start_idx + (k + 1) * block_N, k_end_idx) + T.copy( + K_unpad[k_start_idx + k * block_N : k_copy_end, kv_head_idx, :], + K_shared, + ) + for i, d in T.Parallel(block_N, dim): + if k * block_N + i >= k_current_seqlen: + K_shared[i, d] = 0 + + for i, j in T.Parallel(block_M, block_N): + k_idx = k * block_N + j + boundary_mask = (original_q_idx >= q_current_seqlen) or ( + k_idx >= k_current_seqlen + ) + if is_causal: + row_idx = original_q_idx * block_M + i + cache_len * block_M + orig_row_idx = row_idx // m_block_dim + orig_seqlen_q = ( + (q_current_seqlen + cache_len) * block_M + ) // m_block_dim + compressed_seqlen_q = ( + orig_seqlen_q - kernel_stride + 1 + ) // kernel_stride + offset_row_idx = T.max( + 0, + (orig_row_idx + 1) // kernel_stride + - 1 + + k_current_seqlen + - compressed_seqlen_q, + ) + q_compress_clamped = T.min(k_current_seqlen, offset_row_idx) + causal_mask = k_idx > q_compress_clamped + acc_s[i, j] = T.if_then_else( + boundary_mask or causal_mask, -1e9, 0 + ) + else: + acc_s[i, j] = T.if_then_else(boundary_mask, -1e9, 0) + + T.gemm( + Q_shared, + K_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + + # Normalize and handle NaN/Inf (when logsum is 0 or very small) + for i, j in T.Parallel(block_M, block_N): + normalized = ( + T.exp2(acc_s[i, j] * scale - scores_max[i] * scale) + / logsum[i] + ) + # Handle NaN/Inf: if logsum is very small or result is invalid, set to 0 + acc_s[i, j] = T.if_then_else( + (logsum[i] > 1e-10) + and (normalized >= 0) + and (normalized <= 1e10), + normalized, + T.Cast(accum_dtype, 0.0), + ) + + T.fill(acc_output, 0) + T.reduce_sum(acc_s, acc_output, dim=0) + + # infllmv2 block-based pooling: + # For pool block b, it aggregates k scores in range [b * block_stride - pad_len, b * block_stride - pad_len + num_offs) + # For k_idx, it contributes to pool block b if: + # b * block_stride - pad_len <= k_idx < b * block_stride - pad_len + num_offs + # So: + # start_b = max(0, ceil((k_idx - num_offs + 1 + pad_len) / block_stride)) + # end_b = floor((k_idx + pad_len) / block_stride) + for j in T.Parallel(block_N): + k_idx = k * block_N + j + if ( + original_q_idx < q_current_seqlen + and k_idx < k_current_seqlen + ): + # Calculate which pool blocks this k_idx contributes to + # start_b = ceil((k_idx - num_offs + 1 + pad_len) / block_stride) + # = ceil((k_idx - 5 + 1 + 1) / 4) = ceil((k_idx - 3) / 4) + start_pool = T.max( + 0, + (k_idx - num_offs + 1 + pad_len + block_stride - 1) + // block_stride, + ) + end_pool = T.min( + effective_pooled_k_len - 1, + (k_idx + pad_len) // block_stride, + ) + + pool_block_start = p_block * block_P + pool_block_end = T.min( + (p_block + 1) * block_P, effective_pooled_k_len + ) + + for p_off in T.serial( + num_offs + ): # at most num_offs pool blocks per k + p_idx = start_pool + p_off + if ( + p_idx >= pool_block_start + and p_idx < pool_block_end + and p_idx <= end_pool + ): + local_p_idx = p_idx - pool_block_start + T.atomic_max( + pool_max_shared[local_p_idx], acc_output[j] + ) + T.sync_threads() + + for p_off in T.Parallel(block_P): + p_idx = p_block * block_P + p_off + if ( + p_idx < effective_pooled_k_len + and original_q_idx < q_current_seqlen + ): + off_bq = (original_q_idx + cache_len) // block_size + off_bk = p_idx + + # Match Torch implementation exactly: + # if init_blocks > 0 and off_bk < init_blocks: + # should_mask_inf = True + # elif local_blocks > 0: + # if (off_bq >= off_bk) and (off_bq <= off_bk + local_blocks): + # should_mask_inf = True + is_init_masked = (init_blocks > 0) and (off_bk < init_blocks) + is_local_masked = ( + (local_blocks > 0) + and (off_bq >= off_bk) + and (off_bq <= off_bk + local_blocks) + ) + # Use elif logic: local_blocks check only when not init_masked + is_masked = T.if_then_else( + is_init_masked, 1, T.if_then_else(is_local_masked, 1, 0) + ) + + topk_index_shared[topk + p_off] = p_idx + # Use inf for masked blocks to force selection + # Compare only index sets, not order + topk_value_shared[topk + p_off] = T.if_then_else( + is_masked == 1, + T.Cast("float32", float("inf")), + pool_max_shared[p_off], + ) + T.sync_threads() + + bitonic_sort(topk_index_shared, topk_value_shared) + + for i in T.Parallel(topk): + if original_q_idx < q_current_seqlen: + global_q_idx = q_start_idx + original_q_idx + TopkIndices[kv_head_idx, global_q_idx, i] = topk_index_shared[i] + TopkValues[kv_head_idx, global_q_idx, i] = topk_value_shared[i] + + return main + + +fused_attn_pooling_online_topk_prefill = partial( + _fused_attn_pooling_online_topk, is_causal=True +) +fused_attn_pooling_online_topk_decode = partial( + _fused_attn_pooling_online_topk, + max_seqlen_q_grid=1, + is_causal=False, +) diff --git a/python/sglang/srt/layers/attention/minicpm/sparse_kernels.py b/python/sglang/srt/layers/attention/minicpm/sparse_kernels.py new file mode 100644 index 000000000..878234dc0 --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/sparse_kernels.py @@ -0,0 +1,198 @@ +import triton +import triton.language as tl + + +# TODO. Now only page size == 1 is supported. Consider extend to page size > 1 +@triton.jit +def compress_k_complete_kernel_new( + key_cache_ptr, + token_table_ptr, + cu_new_k_token_nums_ptr, + history_compress_k_token_nums_ptr, + compressed_k_table_ptr, + cu_total_compress_k_token_nums_ptr, + full_compressed_k_ptr, + batch_size, + max_chunks_per_seq, + token_table_cols, + compressed_k_table_cols, + head_num_k: tl.constexpr, + head_dim: tl.constexpr, + kernel_size: tl.constexpr, + kernel_stride: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + max_grid_chunks: tl.constexpr, +): + """ + Single-kernel implementation that fuses k computation, key compression, + key_cache write, and full_compressed_k read for ALL chunks (history + new). + + Grid: (batch_size, min(max_total_chunks, max_grid_chunks), head_num_k) + where max_total_chunks = max_chunks_per_seq + max_history_chunks + - chunk_in_seq in [0, history_chunks_in_seq): process HISTORY chunks + - chunk_in_seq in [history_chunks_in_seq, total_chunks_in_seq): process NEW chunks + + If total_chunks > max_grid_chunks, each thread block loops to handle multiple chunks. + + Each program processes one (batch, chunk_in_seq, head) combination. + """ + batch_idx = tl.program_id(0) + grid_chunk_idx = tl.program_id(1) + head_idx = tl.program_id(2) + + # Total number of chunks this thread block needs to process + chunk_stride = max_grid_chunks + + if batch_idx >= batch_size or head_idx >= head_num_k: + return + + # ==================================================================== + # PHASE 0: Determine chunk type and boundaries + # ==================================================================== + + history_compress = tl.load(history_compress_k_token_nums_ptr + batch_idx) + + # Compute how many NEW chunks this sequence actually has + cu_new_k_start = tl.load(cu_new_k_token_nums_ptr + batch_idx) + cu_new_k_end = tl.load(cu_new_k_token_nums_ptr + batch_idx + 1) + new_k_count = cu_new_k_end - cu_new_k_start + new_chunks_in_seq = tl.where( + new_k_count >= kernel_size, (new_k_count - kernel_size) // kernel_stride + 1, 0 + ) + + # Total chunks = history + new + history_chunks_in_seq = history_compress + total_chunks_in_seq = history_chunks_in_seq + new_chunks_in_seq + + output_start = tl.load(cu_total_compress_k_token_nums_ptr + batch_idx) + + # ==================================================================== + # LOOP: Handle multiple chunks per thread block if needed + # ==================================================================== + + # Iterate over all chunks assigned to this thread block + chunk_in_seq = grid_chunk_idx + + while chunk_in_seq < total_chunks_in_seq: + # Determine if processing history or new chunks + is_history_chunk = chunk_in_seq < history_chunks_in_seq + + if is_history_chunk: + # ==================================================================== + # PHASE 1: Process HISTORY chunks + # ==================================================================== + + # chunk_in_seq in [0, history_compress) -> history chunk index + history_chunk_idx = chunk_in_seq + + global_full_idx = output_start + history_chunk_idx + + # Read from compressed_k_table: indices at y = history_chunk_idx + full_compressed_idx = tl.load( + compressed_k_table_ptr + + batch_idx * compressed_k_table_cols + + history_chunk_idx + ).to(tl.int32) + + head_offset = ( + full_compressed_idx * head_num_k * head_dim + head_idx * head_dim + ) + x = tl.load( + key_cache_ptr + head_offset + tl.arange(0, BLOCK_SIZE), + mask=tl.arange(0, BLOCK_SIZE) < head_dim, + other=0.0, + ) + out_offset = global_full_idx * head_num_k * head_dim + head_idx * head_dim + tl.store( + full_compressed_k_ptr + out_offset + tl.arange(0, BLOCK_SIZE), + x, + mask=tl.arange(0, BLOCK_SIZE) < head_dim, + ) + + else: + # ==================================================================== + # PHASE 2: Process NEW chunks + # ==================================================================== + + # chunk_in_seq in [history_compress, total_chunks_in_seq) -> new chunk index + new_chunk_idx = chunk_in_seq - history_chunks_in_seq + + # Compute y index in token_table for this new chunk + # y = new_chunk_idx * kernel_stride + history_compress * kernel_stride + y = (new_chunk_idx + history_compress) * kernel_stride + + # Use nested if instead of continue (Triton doesn't support continue) + if y < token_table_cols: + # Compute y index in compressed_k_table for new_compressed_k_indices + # y = new_chunk_idx + history_compress + compressed_table_y = new_chunk_idx + history_compress + + if compressed_table_y < compressed_k_table_cols: + # Read new_compressed_k_indices from compressed_k_table + new_compressed_k_indices = tl.load( + compressed_k_table_ptr + + batch_idx * compressed_k_table_cols + + compressed_table_y + ).to(tl.int32) + + # ==================================================================== + # PHASE 3: Perform mean pooling compression on k + # ==================================================================== + + # Accumulate over all tokens in this chunk + acc = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + + for token_offset in range(kernel_size): + # Compute k_indices for this token + token_y = ( + new_chunk_idx * kernel_stride + token_offset + ) + history_compress * kernel_stride + + # Read k_indices from token_table + if token_y < token_table_cols: + token_k_indices = tl.load( + token_table_ptr + batch_idx * token_table_cols + token_y + ).to(tl.int32) + else: + token_k_indices = 0 + + # Load k from key_cache: key_cache[token_k_indices, head_idx, :] + key_base_offset = ( + token_k_indices * head_num_k * head_dim + + head_idx * head_dim + ) + + # Vectorized load of head_dim values + x = tl.load( + key_cache_ptr + key_base_offset + tl.arange(0, BLOCK_SIZE), + mask=tl.arange(0, BLOCK_SIZE) < head_dim, + other=0.0, + ).to(tl.float32) + + acc += x + + # Compute mean over the chunk + acc = acc / kernel_size + + head_offset = ( + new_compressed_k_indices * head_num_k * head_dim + + head_idx * head_dim + ) + tl.store( + key_cache_ptr + head_offset + tl.arange(0, BLOCK_SIZE), + acc, + mask=tl.arange(0, BLOCK_SIZE) < head_dim, + ) + + global_full_idx = output_start + history_compress + new_chunk_idx + out_offset = ( + global_full_idx * head_num_k * head_dim + head_idx * head_dim + ) + tl.store( + full_compressed_k_ptr + out_offset + tl.arange(0, BLOCK_SIZE), + acc, + mask=tl.arange(0, BLOCK_SIZE) < head_dim, + ) + + # Move to next chunk for this thread block + chunk_in_seq += chunk_stride diff --git a/python/sglang/srt/layers/attention/minicpm/sparse_utils.py b/python/sglang/srt/layers/attention/minicpm/sparse_utils.py new file mode 100644 index 000000000..c2402c552 --- /dev/null +++ b/python/sglang/srt/layers/attention/minicpm/sparse_utils.py @@ -0,0 +1,719 @@ +"""Sparse attention utilities for MiniCPM models. + +This module provides sparse attention helpers and utilities for MiniCPM models, +combining both backend-agnostic sparse attention components and kernel utilities. +""" + +from __future__ import annotations + +from itertools import accumulate +from typing import TYPE_CHECKING, Optional + +import msgspec +import torch +import torch.nn.functional as F + +from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionMetadata, +) + +if TYPE_CHECKING: + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +import triton +from sgl_kernel import infllmv2_attn_stage1, max_pooling_1d_varlen + +from sglang.srt.layers.attention.minicpm.sparse_kernels import ( + compress_k_complete_kernel_new, +) +from sglang.srt.model_executor.forward_context import get_token_to_kv_pool + + +def batched_gather(a, lengths_cpu, select): + offsets = [0, *accumulate(map(int, lengths_cpu))] + return torch.cat([a[offsets[i] : offsets[i + 1]] for i in select]) + + +def compress_k_core_new( + full_compressed_k, # output + batch, + key_cache, + token_table, + compressed_k_table, + cu_new_k_token_nums, + history_compress_k_token_nums, + cu_total_compress_k_token_nums, + kernel_size, + kernel_stride, + max_context_length, +): + head_num_k = key_cache.shape[1] + head_dim = key_cache.shape[2] + + # ============================================================================== + # BUFFER ALLOCATION + # ============================================================================== + + # Use provided explicit parameters for buffer allocation + # max_chunks_per_seq is already the maximum possible chunks for any sequence + # given max_context_length, kernel_size, and kernel_stride + max_chunks_per_seq = max(0, (max_context_length - kernel_size) // kernel_stride + 1) + + # ============================================================================== + # Launch kernel for ALL chunks (history + new) + # ============================================================================== + # Grid: (batch, max_chunks_per_seq, head_num_k) + # - chunk_in_seq in [0, history_compress): process HISTORY chunks + # - chunk_in_seq in [history_compress, total_chunks_in_seq): process NEW chunks + # + # max_chunks_per_seq is already the maximum possible chunks for any sequence, + # so it's sufficient for both history and new chunks. + # + # All operations are in a single kernel, CUDA graph compatible. + + # Limit grid size to avoid too many thread blocks + # If max_chunks_per_seq > max_grid_chunks, kernel will loop to handle remaining chunks + MAX_GRID_CHUNKS = 1024 # Adjustable limit for grid dimension + max_grid_chunks = min(max_chunks_per_seq, MAX_GRID_CHUNKS) + + BLOCK_SIZE = triton.next_power_of_2(head_dim) + # Grid size is now limited, kernel uses loop to handle all chunks + grid = (batch, max_grid_chunks, head_num_k) + + compress_k_complete_kernel_new[grid]( + key_cache, + token_table, + cu_new_k_token_nums, + history_compress_k_token_nums, + compressed_k_table, + cu_total_compress_k_token_nums, + full_compressed_k, + batch, + max_chunks_per_seq, + token_table.shape[1], + compressed_k_table.shape[1], + head_num_k, + head_dim, + kernel_size, + kernel_stride, + BLOCK_SIZE, + max_grid_chunks, # Pass the limit to kernel for loop control + ) + + return + + +def get_compress_k_v2( + layer, + forward_batch, + metadata: MiniCPMSparseMetadata, + full_compressed_k1, + full_compressed_k2, + max_context_length, + k1_kernel_size, + k1_kernel_stride, + k2_kernel_size, + k2_kernel_stride, +): + batch = len(forward_batch.req_pool_indices) + key_cache = get_token_to_kv_pool().get_key_buffer(layer.layer_id) + key_cache = key_cache.view(-1, layer.tp_k_head_num, layer.head_dim) + + for full_compressed_k, level, kernel_size, kernel_stride in ( + ( + full_compressed_k1, + metadata.k1, + k1_kernel_size, + k1_kernel_stride, + ), + ( + full_compressed_k2, + metadata.k2, + k2_kernel_size, + k2_kernel_stride, + ), + ): + compress_k_core_new( + full_compressed_k, + batch, + key_cache, + metadata.base.page_table, + level.table, + level.cu_new_token_nums, + level.history_compress_token_nums, + level.cu_total_compress_token_nums, + kernel_size, + kernel_stride, + max_context_length, + ) + + +def allocate_and_compress_keys( + layer, + forward_batch, + metadata: MiniCPMSparseMetadata, + k1_token_nums: int, + k2_token_nums: int, + k1_kernel_size: int, + k1_kernel_stride: int, + k2_kernel_size: int, + k2_kernel_stride: int, + dtype: torch.dtype = torch.bfloat16, + device: torch.device = None, + max_context_length: int = 32768, +): + """Allocate compressed key tensors and run compression. + + Args: + layer: Model layer with head configuration + forward_batch: Forward batch info + metadata: MiniCPM sparse metadata + k1_token_nums: Number of k1 tokens to allocate + k2_token_nums: Number of k2 tokens to allocate + k1_kernel_size: K1 compression window + k1_kernel_stride: K1 compression stride + k2_kernel_size: K2 compression window + k2_kernel_stride: K2 compression stride + dtype: Tensor data type (default: bfloat16) + device: Tensor device (default: layer device) + max_context_length: Maximum context length for the model (default: 32768) + + Returns: + Tuple of (full_compressed_k1, full_compressed_k2) + """ + if device is None: + device = forward_batch.input_ids.device + + full_compressed_k1 = torch.full( + (k1_token_nums, layer.tp_k_head_num, layer.head_dim), + dtype=dtype, + device=device, + fill_value=float("-inf"), + ) + full_compressed_k2 = torch.full( + (k2_token_nums, layer.tp_k_head_num, layer.head_dim), + dtype=dtype, + device=device, + fill_value=float("-inf"), + ) + + get_compress_k_v2( + layer, + forward_batch, + metadata, + full_compressed_k1, + full_compressed_k2, + max_context_length=max_context_length, + k1_kernel_size=k1_kernel_size, + k1_kernel_stride=k1_kernel_stride, + k2_kernel_size=k2_kernel_size, + k2_kernel_stride=k2_kernel_stride, + ) + + return full_compressed_k1, full_compressed_k2 + + +def compressed_attention( + q: torch.Tensor, + k: torch.Tensor, + k2: torch.Tensor, + kernel_stride: int, + block_size: int, + topk: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_k: torch.Tensor, + cu_seqlens_k2: torch.Tensor, + max_seqlen_q: int, + max_context_len: int, + init_blocks: int = 1, + local_blocks: int = 2, + cache_lens: Optional[torch.Tensor] = None, + cu_seqlens_q_adjusted: Optional[torch.Tensor] = None, + max_seqlen_q_adjusted: Optional[int] = None, +) -> torch.Tensor: + """Compressed attention computation for sparse attention. + + Computes attention scores between query and compressed keys (k and k2), + then performs max pooling and selects top-k blocks. + + Args: + q: Query tensor, shape (total_q_len, num_heads, head_dim) + k: Compressed key tensor k1, shape (total_k_len, num_heads, head_dim) + k2: Compressed key tensor k2, shape (total_k_len, num_heads, head_dim) + kernel_stride: Stride of compression kernel + block_size: Size of attention blocks + topk: Number of top blocks to select + cu_seqlens_q: Cumulative sequence lengths for query, shape (batch_size + 1) + cu_seqlens_k: Cumulative sequence lengths for k, shape (batch_size + 1) + cu_seqlens_k2: Cumulative sequence lengths for k2, shape (batch_size + 1) + max_seqlen_q: Maximum sequence length in query + init_blocks: Number of initial blocks to always attend to + local_blocks: Number of local blocks to consider + cache_lens: Cache lengths for each batch (optional) + cu_seqlens_q_adjusted: Adjusted cumulative sequence lengths for query (for stage1 optimization) + max_seqlen_q_adjusted: Adjusted maximum sequence length for query (for stage1 optimization) + + Returns: + Top-k block indices, shape (num_heads, total_q_len, topk) + """ + with torch.no_grad(): + batch_size = cu_seqlens_q.shape[0] - 1 + + is_prefilling = max_seqlen_q > 1 + + if is_prefilling: + if cache_lens is None: + cache_lens = torch.zeros(batch_size, dtype=torch.int32, device=q.device) + + score = infllmv2_attn_stage1( + q.contiguous(), + k.contiguous(), + k2.contiguous(), + cu_seqlens_q=cu_seqlens_q_adjusted, + cu_seqlens_k=cu_seqlens_k, + cu_seqlens_v=cu_seqlens_k2, + max_seqlen_q=max_seqlen_q_adjusted, + max_seqlen_k=max_context_len // kernel_stride, + causal=is_prefilling, + ) + + block_score = max_pooling_1d_varlen( + score.contiguous(), + cu_seqlens_q, + cu_seqlens_k, + cache_lens, + max_seqlen_q, + max_context_len, + local_blocks=local_blocks, + init_blocks=init_blocks, + block_size=block_size, + stride=kernel_stride, + ) + + topk_idx = block_score.topk(topk, dim=-1).indices.sort(-1).values + topk_idx = topk_idx.to(torch.int32) + + return topk_idx + + +def compressed_attention_tilelang( + q: torch.Tensor, + k: torch.Tensor, + block_size: int, + topk: int, + kernel_topk: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_k: torch.Tensor, + cache_lens=None, + fused_kernel=None, + max_cache_len=-1, +) -> torch.Tensor: + """ + 使用 tilelang online topk kernel 计算 compressed attention topk indices + """ + with torch.no_grad(): + batch_size = cu_seqlens_q.shape[0] - 1 + + total_q_len = q.shape[0] + num_kv_heads = k.shape[1] + head_dim = k.shape[2] + + num_heads = q.shape[1] + groups = num_heads // num_kv_heads + q_kernel = q.view(total_q_len, num_kv_heads, groups, head_dim) + q_kernel = ( + q_kernel.transpose(1, 2) + .reshape(total_q_len * groups, num_kv_heads, head_dim) + .contiguous() + ) + + k_kernel = k.contiguous() + + pooled_k_len = (max_cache_len + block_size - 1) // block_size + + assert fused_kernel is not None, "fused_kernel is not initialized" + + # Compute actual output topk (same as original: min(topk, num_blocks)) + output_topk = min(topk, pooled_k_len) + + # Allocate output tensors + topk_indices = torch.full( + (num_kv_heads, total_q_len, kernel_topk), + -1, + dtype=torch.int32, + device=q.device, + ) + topk_values = torch.full( + (num_kv_heads, total_q_len, kernel_topk), + float("-inf"), + dtype=torch.float32, + device=q.device, + ) + + if cache_lens is None: + cache_lens_tensor = torch.zeros( + batch_size, dtype=torch.int32, device=q.device + ) + else: + cache_lens_tensor = cache_lens.to(torch.int32) + + fused_kernel( + q_kernel, + k_kernel, + cu_seqlens_q, + cu_seqlens_k, + cache_lens_tensor, + topk_indices, + topk_values, + ) + + # Note: q_idx masking is handled inside the kernel via causal_mask + # which sets scores to -1e9 for K blocks beyond the causal boundary. + # These blocks won't be selected in topk due to their low scores. + + # Sort with -1 values at the end (match original behavior) + # Replace -1 with large value, sort, then replace back + large_val = pooled_k_len + 1000 # Any value larger than max valid index + topk_for_sort = topk_indices.clone() + topk_for_sort[topk_for_sort == -1] = large_val + topk_idx = topk_for_sort.sort(-1).values + topk_idx[topk_idx == large_val] = -1 + + # Truncate to output_topk (same as original: min(topk, num_blocks)) + topk_idx = topk_idx[:, :, :output_topk].contiguous() + + return topk_idx + + +class CompressionLevelMetadata(msgspec.Struct): + """Metadata for a single compression level (k1 or k2). + + This struct groups all metadata fields for one compression level, + reducing duplication and making the code more maintainable. + """ + + # Cumulative sequence lengths for compressed cache + cu_seqlens: Optional[torch.Tensor] = None + cu_seqlens_cpu: Optional[list[int]] = None + + # Token mapping table (request pool indices -> compressed cache tokens) + table: Optional[torch.Tensor] = None + + # Compressed cache metadata + history_compress_token_nums: Optional[torch.Tensor] = None + cu_new_token_nums: Optional[torch.Tensor] = None + cu_total_compress_token_nums: Optional[torch.Tensor] = None + + +class MiniCPMSparseMetadata(msgspec.Struct): + base: FlashAttentionMetadata + k1: Optional[CompressionLevelMetadata] = None + k2: Optional[CompressionLevelMetadata] = None + sparse_bs_list: Optional[list[int]] = None + sparse_idx: Optional[list[int]] = None + dense_layout: Optional[list[tuple[int, int, int, int]]] = None + seqlen_k_sparse_bs_tensor: Optional[torch.Tensor] = None + token_to_bs: Optional[torch.Tensor] = None + token_pos_in_bs: Optional[torch.Tensor] = None + sparse_page_table: Optional[torch.Tensor] = None + sparse_cache_seqlens_int32: Optional[torch.Tensor] = None + sparse_cu_seqlens_q: Optional[torch.Tensor] = None + sparse_cu_seqlens_k: Optional[torch.Tensor] = None + sparse_max_seq_len_q: int = 1 + cache_seqlens_int32_stage1: Optional[torch.Tensor] = None + cu_seqlens_q_adjusted: Optional[torch.Tensor] = None + max_seqlen_q_adjusted: int = 1 + topk_cu_seqlens_q: Optional[torch.Tensor] = None + topk_cu_seqlens_k: Optional[torch.Tensor] = None + topk_max_seqlen_q: int = 1 + topk_max_seqlen_k: int = 1 + + +def _compute_single_compression_metadata( + seq_lens_cpu: torch.Tensor, + token_nums: torch.Tensor, + history_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_sparse_token: torch.Tensor, + kernel_size: int, + kernel_stride: int, +) -> CompressionLevelMetadata: + seqlen_cpu = torch.clamp( + (seq_lens_cpu - kernel_size) // kernel_stride + 1, + min=0, + ) + + cu_seqlens_cpu = F.pad( + torch.cumsum(seqlen_cpu, dim=0, dtype=torch.int32), (1, 0) + ).tolist() + cu_seqlens = F.pad( + torch.cumsum(seqlen_cpu.to(device=token_nums.device), dim=0, dtype=torch.int32), + (1, 0), + ) + token_table = req_to_sparse_token[req_pool_indices] + history_compress_token_nums = torch.clamp( + (history_lens - kernel_size) // kernel_stride + 1, + min=0, + ) + new_token_nums = token_nums - history_compress_token_nums * kernel_stride + cu_new_token_nums = F.pad( + torch.cumsum(new_token_nums, dim=0, dtype=torch.int32), (1, 0) + ) + new_compress_token_nums = torch.clamp( + (new_token_nums - kernel_size) // kernel_stride + 1, + min=0, + ) + total_compress_token_nums = history_compress_token_nums + new_compress_token_nums + cu_total_compress_token_nums = F.pad( + torch.cumsum(total_compress_token_nums, dim=0, dtype=torch.int32), (1, 0) + ) + + return CompressionLevelMetadata( + cu_seqlens=cu_seqlens, + cu_seqlens_cpu=cu_seqlens_cpu, + table=token_table, + history_compress_token_nums=history_compress_token_nums, + cu_new_token_nums=cu_new_token_nums, + cu_total_compress_token_nums=cu_total_compress_token_nums, + ) + + +def _build_k1_k2_compression_metadata( + forward_batch: ForwardBatch, + base_metadata: FlashAttentionMetadata, + req_to_sparse_k1_token: torch.Tensor, + req_to_sparse_k2_token: torch.Tensor, + k1_kernel_size: int, + k1_kernel_stride: int, + k2_kernel_size: int, + k2_kernel_stride: int, + cu_seqlens_q: torch.Tensor, +) -> tuple[CompressionLevelMetadata, CompressionLevelMetadata]: + bs = forward_batch.batch_size + seq_lens_cpu = torch.as_tensor( + forward_batch.seq_lens_cpu, + dtype=base_metadata.cu_seqlens_q.dtype, + device="cpu", + ) + token_nums = ( + base_metadata.cu_seqlens_k[1 : bs + 1] - base_metadata.cu_seqlens_k[:bs] + ) + input_lens = cu_seqlens_q[1 : bs + 1] - cu_seqlens_q[:bs] + history_lens = token_nums - input_lens + + return tuple( + _compute_single_compression_metadata( + seq_lens_cpu, + token_nums, + history_lens, + forward_batch.req_pool_indices, + req_to_sparse_token, + kernel_size, + kernel_stride, + ) + for req_to_sparse_token, kernel_size, kernel_stride in ( + (req_to_sparse_k1_token, k1_kernel_size, k1_kernel_stride), + (req_to_sparse_k2_token, k2_kernel_size, k2_kernel_stride), + ) + ) + + +def _get_sparse_cache_lens( + seq_lens: torch.Tensor, + sparse_capacity: int, + block_size: int, +) -> torch.Tensor: + remainder = seq_lens % block_size + sparse_lens = torch.where( + remainder == 0, + sparse_capacity, + sparse_capacity - block_size + remainder, + ) + return torch.where(seq_lens <= sparse_capacity, seq_lens, sparse_lens) + + +def _plan_sparse_prefill( + forward_batch: ForwardBatch, + metadata: MiniCPMSparseMetadata, + head_group_num: int, + heads_per_group: int, + dense_len: int, + sparse_topk: int, + block_size: int, +) -> None: + device = metadata.base.cu_seqlens_q.device + sparse_capacity = sparse_topk * block_size + sparse_bs_list = [] + sparse_idx = [] + dense_layout = [] + row_q_lens = [] + sparse_cache_seqlens = [] + token_to_bs = [] + token_pos_in_bs = [] + sparse_q_lens = [] + sparse_k_lens = [] + dense_q_lens = [] + max_sparse_cache_len = 0 + query_group_start = 0 + + for batch_idx in range(forward_batch.batch_size): + seq_len = int(forward_batch.seq_lens_cpu[batch_idx]) + query_len = int(forward_batch.extend_seq_lens_cpu[batch_idx]) + prefix_len = int(forward_batch.extend_prefix_lens_cpu[batch_idx]) + row_start = len(row_q_lens) + if seq_len >= dense_len: + sparse_batch_idx = len(sparse_bs_list) + sparse_bs_list.append(batch_idx) + sparse_q_lens.append(query_len) + sparse_k_lens.append(seq_len) + sparse_idx.extend(range(row_start, row_start + query_len * head_group_num)) + row_q_lens.extend([1] * (query_len * head_group_num)) + token_to_bs.extend([sparse_batch_idx] * query_len) + token_pos_in_bs.extend(range(prefix_len + 1, prefix_len + query_len + 1)) + token_seq_lens = torch.arange( + prefix_len + 1, + prefix_len + query_len + 1, + dtype=torch.int32, + ) + sparse_cache_seqlens.extend( + _get_sparse_cache_lens(token_seq_lens, sparse_capacity, block_size) + .repeat_interleave(head_group_num) + .tolist() + ) + max_sparse_cache_len = max(max_sparse_cache_len, sparse_capacity) + else: + dense_layout.append((batch_idx, row_start, query_group_start, query_len)) + dense_q_lens.append(query_len) + row_q_lens.extend([query_len] * head_group_num) + sparse_cache_seqlens.extend([seq_len] * head_group_num) + max_sparse_cache_len = max(max_sparse_cache_len, seq_len) + query_group_start += query_len * head_group_num + + metadata.sparse_bs_list = sparse_bs_list + metadata.sparse_idx = sparse_idx + metadata.dense_layout = dense_layout + metadata.token_to_bs = torch.tensor(token_to_bs, dtype=torch.int32, device=device) + metadata.token_pos_in_bs = torch.tensor( + token_pos_in_bs, dtype=torch.int32, device=device + ) + metadata.seqlen_k_sparse_bs_tensor = torch.tensor( + sparse_k_lens, dtype=torch.int32, device=device + ) + metadata.sparse_page_table = torch.zeros( + (len(row_q_lens), max_sparse_cache_len), + dtype=metadata.base.page_table.dtype, + device=metadata.base.page_table.device, + ) + row_q_lens_tensor = torch.tensor( + row_q_lens, dtype=metadata.base.cu_seqlens_q.dtype, device=device + ) + metadata.sparse_cu_seqlens_q = F.pad( + torch.cumsum(row_q_lens_tensor, dim=0, dtype=torch.int32), (1, 0) + ) + metadata.sparse_max_seq_len_q = max(dense_q_lens, default=1) + metadata.sparse_cache_seqlens_int32 = torch.tensor( + sparse_cache_seqlens, + dtype=torch.int32, + device=device, + ) + metadata.sparse_cu_seqlens_k = F.pad( + torch.cumsum(metadata.sparse_cache_seqlens_int32, dim=0, dtype=torch.int32), + (1, 0), + ) + metadata.cache_seqlens_int32_stage1 = ( + metadata.base.cache_seqlens_int32[sparse_bs_list] - 1 + ) + + if sparse_bs_list: + sparse_q_lens_tensor = torch.tensor( + sparse_q_lens, dtype=torch.int32, device=device + ) + metadata.topk_cu_seqlens_q = F.pad( + torch.cumsum(sparse_q_lens_tensor, dim=0, dtype=torch.int32), (1, 0) + ) + metadata.topk_cu_seqlens_k = F.pad( + torch.cumsum(metadata.seqlen_k_sparse_bs_tensor, dim=0, dtype=torch.int32), + (1, 0), + ) + metadata.topk_max_seqlen_q = max(sparse_q_lens) + metadata.topk_max_seqlen_k = max(sparse_k_lens) + metadata.cu_seqlens_q_adjusted = metadata.topk_cu_seqlens_q * heads_per_group + metadata.max_seqlen_q_adjusted = metadata.topk_max_seqlen_q * heads_per_group + else: + metadata.cu_seqlens_q_adjusted = metadata.base.cu_seqlens_q * heads_per_group + metadata.max_seqlen_q_adjusted = metadata.base.max_seq_len_q * heads_per_group + + +def _plan_sparse_decode( + forward_batch: ForwardBatch, + metadata: MiniCPMSparseMetadata, + head_group_num: int, + dense_len: int, + sparse_topk: int, + block_size: int, +) -> None: + base_metadata = metadata.base + bs = forward_batch.batch_size + cache_seqlens = base_metadata.cache_seqlens_int32 + page_table = base_metadata.page_table + seq_lens_cpu = torch.as_tensor( + forward_batch.seq_lens_cpu, dtype=cache_seqlens.dtype, device="cpu" + ) + sparse_capacity = sparse_topk * block_size + cache_lens_cpu = torch.where( + seq_lens_cpu >= dense_len, + _get_sparse_cache_lens(seq_lens_cpu, sparse_capacity, block_size), + seq_lens_cpu, + ) + sparse_mask_cpu = seq_lens_cpu >= dense_len + sparse_bs_list = sparse_mask_cpu.nonzero().flatten().tolist() + dense_bs_list = (~sparse_mask_cpu).nonzero().flatten().tolist() + sparse_idx = [ + row + for batch_idx in sparse_bs_list + for row in range(batch_idx * head_group_num, (batch_idx + 1) * head_group_num) + ] + max_sparse_cache_len = max( + int(cache_lens_cpu.max()), + sparse_capacity if sparse_bs_list else 0, + ) + sparse_cache_seqlens_cpu = cache_lens_cpu.repeat_interleave(head_group_num) + + sparse_cache_seqlens_int32 = sparse_cache_seqlens_cpu.to( + device=cache_seqlens.device + ) + sparse_cu_seqlens_k = F.pad( + torch.cumsum(sparse_cache_seqlens_int32, dim=0, dtype=torch.int32), (1, 0) + ) + sparse_cu_seqlens_q = torch.arange( + 0, + bs * head_group_num + 1, + dtype=torch.int32, + device=base_metadata.cu_seqlens_q.device, + ) + token_to_bs = torch.arange( + 0, len(sparse_bs_list), dtype=torch.int32, device=page_table.device + ) + sparse_page_table = torch.zeros( + (head_group_num * bs, max_sparse_cache_len), + dtype=page_table.dtype, + device=page_table.device, + ) + + metadata.sparse_cache_seqlens_int32 = sparse_cache_seqlens_int32 + metadata.sparse_cu_seqlens_k = sparse_cu_seqlens_k + metadata.sparse_cu_seqlens_q = sparse_cu_seqlens_q + metadata.sparse_page_table = sparse_page_table + metadata.sparse_bs_list = sparse_bs_list + metadata.sparse_idx = sparse_idx + metadata.dense_layout = [ + (batch_idx, batch_idx * head_group_num, batch_idx * head_group_num, 1) + for batch_idx in dense_bs_list + ] + metadata.token_to_bs = token_to_bs + metadata.topk_cu_seqlens_q = torch.arange( + 0, + len(sparse_bs_list) + 1, + dtype=torch.int32, + device=base_metadata.cu_seqlens_q.device, + ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index a715e67b5..41ed82a87 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -4386,6 +4386,7 @@ class Scheduler( self.tree_cache.reset() self.req_to_token_pool.clear() self.token_to_kv_pool_allocator.clear() + self.req_to_token_pool.reset_aux_cache_allocator() self.grammar_manager.clear() self.metrics_reporter.reset_metrics() diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index 5c4b06a61..2d42bb593 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -98,7 +98,9 @@ class SchedulerInvariantChecker: else: protected = self.tree_cache.protected_size() session_held = self.pool_stats_observer.session_held_tokens() - total = self.token_to_kv_pool_allocator.size + total = self.req_to_token_pool.schedulable_token_capacity( + self.token_to_kv_pool_allocator.size + ) else: protected = self.tree_cache.protected_size() session_held = self.pool_stats_observer.session_held_tokens() @@ -164,12 +166,20 @@ class SchedulerInvariantChecker: return leak, msg free_full_pages = set(free_pages.tolist() + release_pages.tolist()) cached_full_pages = set(self.tree_cache.all_values_flatten().tolist()) - expected_full_pages = set( - range(1, self.token_to_kv_pool_allocator.size + 1) - ) - leaked_full_pages = ( - expected_full_pages - free_full_pages - cached_full_pages - ) + full_page_msg = "" + if ( + self.req_to_token_pool.schedulable_token_capacity( + self.token_to_kv_pool_allocator.size + ) + == self.token_to_kv_pool_allocator.size + ): + expected_full_pages = set( + range(1, self.token_to_kv_pool_allocator.size + 1) + ) + leaked_full_pages = ( + expected_full_pages - free_full_pages - cached_full_pages + ) + full_page_msg = f", leaked_full_pages={leaked_full_pages or None}" mamba_allocator = self.req_to_token_pool.mamba_allocator free_mamba_pages = set(mamba_allocator.free_slots.tolist()) cached_mamba_pages = set( @@ -179,10 +189,8 @@ class SchedulerInvariantChecker: leaked_mamba_pages = ( expected_mamba_pages - free_mamba_pages - cached_mamba_pages ) - msg += ( - f", leaked_full_pages={leaked_full_pages or None}" - f", leaked_mamba_pages={leaked_mamba_pages or None}" - ) + msg += full_page_msg + msg += f", leaked_mamba_pages={leaked_mamba_pages or None}" return leak, msg def _check_mamba_pool_with_int8(self, ps: PoolStats, ckpt_pool) -> Tuple[bool, str]: diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index 087e266fc..dd9cdc7d7 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -261,13 +261,14 @@ class SchedulerPoolStatsObserver: if (is_mamba_radix_cache and not has_int8_ckpt) else 0 ) - full_num_used = self.token_to_kv_pool_allocator.size - ( - full_available_size + full_evictable_size + full_capacity = self.req_to_token_pool.schedulable_token_capacity( + self.token_to_kv_pool_allocator.size ) + full_num_used = full_capacity - (full_available_size + full_evictable_size) mamba_num_used = self.req_to_token_pool.mamba_pool.size - ( mamba_available_size + mamba_evictable_size ) - full_token_usage = full_num_used / self.token_to_kv_pool_allocator.size + full_token_usage = full_num_used / full_capacity mamba_usage = mamba_num_used / self.req_to_token_pool.mamba_pool.size return PoolStats( diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 6dfcb618b..9db800ccc 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -536,7 +536,9 @@ class TpModelWorker(BaseTpWorker): - 1, ) return ( - self.model_runner.max_total_num_tokens, + self.model_runner.req_to_token_pool.schedulable_token_capacity( + self.model_runner.max_total_num_tokens + ), get_schedule().max_prefill_tokens, self.model_runner.max_running_requests, get_schedule().max_queued_requests, diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index fc5428a34..2b8105e3a 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -359,6 +359,14 @@ def alloc_for_extend( prefix_tensors, batch.req_to_token_pool, ) + try: + batch.req_to_token_pool.alloc_aux_to_lengths( + req_pool_indices_cpu=req_pool_indices_cpu, + target_seq_lens_cpu=batch.seq_lens_cpu, + ) + except Exception: + batch.tree_cache.token_to_kv_pool_allocator.free(out_cache_loc) + raise # DSV4-NPU hook: no-op on non-DSV4 paths. if _is_npu: @@ -559,6 +567,15 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor: token_per_req, ) + try: + batch.req_to_token_pool.alloc_aux_to_lengths( + req_pool_indices_cpu=batch.req_pool_indices_cpu, + target_seq_lens_cpu=batch.seq_lens_cpu + token_per_req, + ) + except Exception: + batch.tree_cache.token_to_kv_pool_allocator.free(out_cache_loc) + raise + for req in batch.reqs: req.kv.kv_allocated_len += token_per_req diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 43f4f26d7..9cb0fb363 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -281,6 +281,7 @@ class ReqToTokenPool: ) self.free_slots = list(range(1, self._alloc_size)) self.req_generation = torch.zeros(self._alloc_size, dtype=torch.int64) + self._aux_cache: Any = None def write(self, indices, values): self.req_to_token[indices] = values @@ -324,12 +325,41 @@ class ReqToTokenPool: def free(self, req: Req): assert req.req_pool_idx is not None, "request must have req_pool_idx" + if self._aux_cache is not None: + self._aux_cache.free(req.req_pool_idx) self.free_slots.append(req.req_pool_idx) req.req_pool_idx = None def clear(self): self.free_slots = list(range(1, self._alloc_size)) self.req_generation.zero_() + if self._aux_cache is not None: + self._aux_cache.clear() + + def attach_aux_cache(self, aux_cache: Any) -> None: + assert self._aux_cache is None + self._aux_cache = aux_cache + + def reset_aux_cache_allocator(self) -> None: + if self._aux_cache is not None: + self._aux_cache.reset_allocator() + + def schedulable_token_capacity(self, physical_capacity: int) -> int: + if self._aux_cache is None: + return physical_capacity + return self._aux_cache.dense_capacity + + def alloc_aux_to_lengths( + self, + *, + req_pool_indices_cpu: torch.Tensor, + target_seq_lens_cpu: torch.Tensor, + ) -> None: + if self._aux_cache is not None: + self._aux_cache.alloc_to_lengths( + req_pool_indices_cpu=req_pool_indices_cpu, + target_seq_lens_cpu=target_seq_lens_cpu, + ) class MambaPool: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index e69d4b9ad..9ad07a41d 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1298,9 +1298,12 @@ class ModelRunner: def effective_max_total_num_tokens(self): """Return the max token pool size considering hybrid swa settings.""" if self.is_hybrid_swa: - return self.full_max_total_num_tokens or self.swa_max_total_num_tokens + capacity = self.full_max_total_num_tokens or self.swa_max_total_num_tokens else: - return self.max_total_num_tokens + capacity = self.max_total_num_tokens + if (req_to_token_pool := getattr(self, "req_to_token_pool", None)) is not None: + return req_to_token_pool.schedulable_token_capacity(capacity) + return capacity @property def max_token_pool_size(self): diff --git a/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py b/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py index be3c3d157..eb979f87f 100644 --- a/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py +++ b/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py @@ -96,6 +96,7 @@ def compute_post_capture_kv_resize( ) pool.finalize_backing(config) model_runner.token_to_kv_pool_allocator.resize(config) + model_runner.req_to_token_pool.reset_aux_cache_allocator() capped_max_running_requests = None if model_runner.max_running_requests is not None: diff --git a/python/sglang/srt/models/minicpm.py b/python/sglang/srt/models/minicpm.py index 63efbf1f2..3ec0fff8f 100644 --- a/python/sglang/srt/models/minicpm.py +++ b/python/sglang/srt/models/minicpm.py @@ -17,11 +17,14 @@ import math from typing import Any, Dict, Iterable, Optional, Tuple import torch +import torch.nn.functional as F from torch import nn +from sglang.srt.configs.minicpm import MiniCPMHybridConfig from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( + ColumnParallelLinear, MergedColumnParallelLinear, QKVParallelLinear, RowParallelLinear, @@ -35,9 +38,12 @@ from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.model_loader.weight_utils import ( + default_weight_loader, + sharded_weight_loader, +) from sglang.srt.runtime_context import get_parallel -from sglang.srt.utils import add_prefix +from sglang.srt.utils import add_prefix, set_weight_attrs from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -85,11 +91,15 @@ class MiniCPMAttention(nn.Module): hidden_size: int, num_heads: int, num_kv_heads: int, + head_dim: Optional[int] = None, layer_id: int = 0, rope_theta: float = 10000, rope_scaling: Optional[Dict[str, Any]] = None, max_position_embeddings: int = 8192, quant_config: Optional[QuantizationConfig] = None, + attn_use_rope: bool = True, + use_output_gate: bool = False, + attention_bias: bool = False, prefix: str = "", ) -> None: super().__init__() @@ -108,37 +118,42 @@ class MiniCPMAttention(nn.Module): # the KV heads across multiple tensor parallel GPUs. assert tp_size % self.total_num_kv_heads == 0 self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) - self.head_dim = hidden_size // self.total_num_heads + self.head_dim = ( + head_dim if head_dim is not None else hidden_size // self.total_num_heads + ) self.q_size = self.num_heads * self.head_dim self.kv_size = self.num_kv_heads * self.head_dim self.scaling = self.head_dim**-0.5 self.rope_theta = rope_theta self.max_position_embeddings = max_position_embeddings + self.attn_use_rope = attn_use_rope + self.use_output_gate = use_output_gate self.qkv_proj = QKVParallelLinear( hidden_size, self.head_dim, self.total_num_heads, self.total_num_kv_heads, - bias=False, + bias=attention_bias, quant_config=quant_config, prefix=add_prefix("qkv_proj", prefix), ) self.o_proj = RowParallelLinear( self.total_num_heads * self.head_dim, hidden_size, - bias=False, + bias=attention_bias, quant_config=quant_config, prefix=add_prefix("o_proj", prefix), ) - self.rotary_emb = get_rope( - self.head_dim, - rotary_dim=self.head_dim, - max_position=max_position_embeddings, - base=rope_theta, - rope_scaling=rope_scaling, - ) + if self.attn_use_rope: + self.rotary_emb = get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=rope_theta, + rope_scaling=rope_scaling, + ) self.attn = RadixAttention( self.num_heads, self.head_dim, @@ -149,21 +164,194 @@ class MiniCPMAttention(nn.Module): prefix=add_prefix("attn", prefix), ) + if self.use_output_gate: + self.o_gate = ColumnParallelLinear( + hidden_size, + self.total_num_heads * self.head_dim, + bias=attention_bias, + quant_config=quant_config, + prefix=add_prefix("o_gate", prefix), + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + + if self.attn_use_rope: + orig_dtype = q.dtype + q, k = q.float(), k.float() + q, k = self.rotary_emb(positions, q, k) + q, k = q.to(orig_dtype), k.to(orig_dtype) + + attn_output = self.attn(q, k, v, forward_batch) + + if self.use_output_gate: + o_gate_output, _ = self.o_gate(hidden_states) + attn_output = attn_output * F.sigmoid(o_gate_output) + + output, _ = self.o_proj(attn_output) + return output + + +class MiniCPMLightningMixer(nn.Module): + """Lightning attention mixer backed by the shared linear-attention backend. + + This is a wrapper that prepares inputs for the backend and handles + the QKV projection, normalization, RoPE, and output processing, + while delegating the recurrent computation through RadixAttention. + """ + + def __init__( + self, + hidden_size: int, + num_heads: int, + num_kv_heads: int, + head_dim: int, + layer_id: int = 0, + rope_theta: float = 10000, + rope_scaling: Optional[Dict[str, Any]] = None, + max_position_embeddings: int = 8192, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + use_rope: bool = True, + use_output_gate: bool = False, + attention_bias: bool = False, + rms_norm_eps: float = 1e-6, + use_output_norm: bool = False, + qk_norm: bool = True, + scale: str | float = "1/sqrt(d)", + ) -> None: + super().__init__() + self.hidden_size = hidden_size + tp_size = get_parallel().tp_size + self.total_num_heads = num_heads + assert self.total_num_heads % tp_size == 0 + self.num_heads = self.total_num_heads // tp_size + self.total_num_kv_heads = num_kv_heads + if self.total_num_kv_heads >= tp_size: + assert self.total_num_kv_heads % tp_size == 0 + else: + assert tp_size % self.total_num_kv_heads == 0 + self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) + self.head_dim = head_dim + if scale == "1/sqrt(d)": + scaling = self.head_dim ** (-0.5) + elif scale == "1/d": + scaling = self.head_dim ** (-1.0) + elif isinstance(scale, (int, float)): + scaling = float(scale) + else: + raise ValueError(f"Unsupported lightning scale: {scale}") + self.use_output_gate = use_output_gate + self.attention_bias = attention_bias + self.rms_norm_eps = rms_norm_eps + self.use_rope = use_rope + self.qk_norm = qk_norm + self.use_output_norm = use_output_norm + + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim + + self.qkv_proj = QKVParallelLinear( + hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=self.attention_bias, + quant_config=quant_config, + prefix=add_prefix("qkv_proj", prefix), + ) + + self.o_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + hidden_size, + bias=self.attention_bias, + quant_config=quant_config, + prefix=add_prefix("o_proj", prefix), + ) + + if self.use_output_norm: + self.o_norm = RMSNorm(self.num_heads * self.head_dim, eps=self.rms_norm_eps) + set_weight_attrs( + self.o_norm.weight, {"weight_loader": sharded_weight_loader(0)} + ) + + if self.use_output_gate: + self.z_proj = ColumnParallelLinear( + self.hidden_size, + self.total_num_heads * self.head_dim, + bias=self.attention_bias, + quant_config=quant_config, + prefix=add_prefix("z_proj", prefix), + ) + + if self.qk_norm: + self.q_norm = RMSNorm(self.head_dim, eps=self.rms_norm_eps) + self.k_norm = RMSNorm(self.head_dim, eps=self.rms_norm_eps) + + if self.use_rope: + self.rotary_emb = get_rope( + self.head_dim, + rotary_dim=self.head_dim, + max_position=max_position_embeddings, + base=rope_theta, + rope_scaling=rope_scaling, + ) + + self.attn = RadixAttention( + self.num_heads, + self.head_dim, + scaling, + num_kv_heads=self.num_kv_heads, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("attn", prefix), + ) + def forward( self, positions: torch.Tensor, hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: + if forward_batch.forward_mode.is_idle(): + return hidden_states.new_empty(hidden_states.shape[0], self.hidden_size) + qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - orig_dtype = q.dtype - q, k = q.float(), k.float() - q, k = self.rotary_emb(positions, q, k) - q, k = q.to(orig_dtype), k.to(orig_dtype) - attn_output = self.attn(q, k, v, forward_batch) - output, _ = self.o_proj(attn_output) - return output + + if self.qk_norm: + q = self.q_norm(q.reshape(-1, self.head_dim)) + k = self.k_norm(k.reshape(-1, self.head_dim)) + + if self.use_rope: + q = q.reshape(-1, self.num_heads * self.head_dim) + k = k.reshape(-1, self.num_kv_heads * self.head_dim) + orig_dtype = q.dtype + q, k = q.float(), k.float() + q, k = self.rotary_emb(positions, q, k) + q, k = q.to(orig_dtype), k.to(orig_dtype) + + q = q.reshape(-1, self.num_heads, self.head_dim) + k = k.reshape(-1, self.num_kv_heads, self.head_dim) + v = v.reshape(-1, self.num_kv_heads, self.head_dim) + + o = self.attn(q, k, v, forward_batch) + + if self.use_output_norm: + o = self.o_norm(o) + + if self.use_output_gate: + z, _ = self.z_proj(hidden_states) + o = o * F.sigmoid(z) + + y, _ = self.o_proj(o) + return y class MiniCPMDecoderLayer(nn.Module): @@ -176,20 +364,59 @@ class MiniCPMDecoderLayer(nn.Module): ) -> None: super().__init__() self.config = config + self.layer_id = layer_id self.hidden_size = config.hidden_size + if isinstance(config, MiniCPMHybridConfig): + self.mixer_type = config.mixer_types[layer_id] + attn_use_rope = config.attn_use_rope + attn_use_output_gate = config.attn_use_output_gate + attention_bias = config.attention_bias + else: + self.mixer_type = "minicpm4" + attn_use_rope = True + attn_use_output_gate = False + attention_bias = False + rope_theta, rope_scaling = get_rope_config(config) max_position_embeddings = getattr(config, "max_position_embeddings", 8192) - self.self_attn = MiniCPMAttention( - hidden_size=self.hidden_size, - num_heads=config.num_attention_heads, - num_kv_heads=config.num_key_value_heads, - layer_id=layer_id, - rope_theta=rope_theta, - rope_scaling=rope_scaling, - max_position_embeddings=max_position_embeddings, - quant_config=quant_config, - prefix=add_prefix("self_attn", prefix), - ) + if self.mixer_type == "minicpm4": + self.self_attn = MiniCPMAttention( + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + num_kv_heads=config.num_key_value_heads, + head_dim=getattr(config, "head_dim", None), + layer_id=layer_id, + rope_theta=rope_theta, + rope_scaling=rope_scaling, + max_position_embeddings=max_position_embeddings, + quant_config=quant_config, + attn_use_rope=attn_use_rope, + use_output_gate=attn_use_output_gate, + attention_bias=attention_bias, + prefix=add_prefix("self_attn", prefix), + ) + elif self.mixer_type == "lightning-attn": + self.self_attn = MiniCPMLightningMixer( + hidden_size=self.hidden_size, + num_heads=config.lightning_nh, + num_kv_heads=config.lightning_nkv, + head_dim=config.lightning_head_dim, + layer_id=layer_id, + rope_theta=rope_theta, + rope_scaling=rope_scaling, + max_position_embeddings=max_position_embeddings, + quant_config=quant_config, + use_rope=config.lightning_use_rope, + use_output_gate=config.use_output_gate, + attention_bias=config.attention_bias, + rms_norm_eps=config.rms_norm_eps, + use_output_norm=config.use_output_norm, + qk_norm=config.qk_norm, + scale=config.lightning_scale, + prefix=add_prefix("self_attn", prefix), + ) + else: + raise ValueError(f"Unsupported mixer type: {self.mixer_type}") self.mlp = MiniCPMMLP( hidden_size=self.hidden_size, intermediate_size=config.intermediate_size, @@ -287,7 +514,7 @@ class MiniCPMModel(nn.Module): return hidden_states -class MiniCPMForCausalLM(nn.Module): +class MiniCPMSALAForCausalLM(nn.Module): def __init__( self, config, @@ -396,4 +623,8 @@ class MiniCPMForCausalLM(nn.Module): weight_loader(param, loaded_weight) -EntryClass = MiniCPMForCausalLM +class MiniCPMForCausalLM(MiniCPMSALAForCausalLM): + """Alias for MiniCPM checkpoints whose config uses the HF architecture name.""" + + +EntryClass = [MiniCPMSALAForCausalLM, MiniCPMForCausalLM] diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a63a403e1..80136a2a7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -201,6 +201,8 @@ ATTENTION_BACKEND_CHOICES = [ "trtllm_mha", "dual_chunk_flash_attn", "hpc_ops", # HPC-Ops (https://github.com/Tencent/hpc-ops), Hopper (SM90) only, requires --page-size 64 + "minicpm_flashattn", + "minicpm_flashinfer", # AMD specific "aiter", "wave", @@ -969,7 +971,12 @@ class ServerArgs: NS("schedule"), ] = False disable_radix_cache: A[ - bool, "Disable RadixAttention for prefix caching.", NS("memory") + bool, + Arg( + help="Disable RadixAttention for prefix caching.", + resolvable=True, + ), + NS("memory"), ] = False enable_page_major_kv_layout: A[ bool, diff --git a/python/sglang/srt/session/streaming_session.py b/python/sglang/srt/session/streaming_session.py index 288cd3de5..bdbf38d75 100644 --- a/python/sglang/srt/session/streaming_session.py +++ b/python/sglang/srt/session/streaming_session.py @@ -465,7 +465,7 @@ class StreamingSession(BasePrefixCache): slot.req_pool_idx, start:end ] self.token_to_kv_pool_allocator.free(kv_indices) - self.req_to_token_pool.free_slots.append(slot.req_pool_idx) + self.req_to_token_pool.free(slot) self._free_slot_mamba(slot) diff --git a/python/sglang/srt/utils/hf_transformers/common.py b/python/sglang/srt/utils/hf_transformers/common.py index 21f5ca8af..c5f346786 100644 --- a/python/sglang/srt/utils/hf_transformers/common.py +++ b/python/sglang/srt/utils/hf_transformers/common.py @@ -49,6 +49,7 @@ from sglang.srt.configs import ( LagunaConfig, LocateAnythingConfig, LongcatFlashConfig, + MiniCPMHybridConfig, MiniCPMV4_6Config, MiniCPMV4_6VisionConfig, MiniMaxM3VLConfig, @@ -131,6 +132,7 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = { KimiK25Config, Step3p5Config, Step3p7Config, + MiniCPMHybridConfig, MiniCPMV4_6Config, MiniCPMV4_6VisionConfig, InklingModelConfig, diff --git a/test/registered/kernels/benchmark/bench_get_block_table.py b/test/registered/kernels/benchmark/bench_get_block_table.py new file mode 100644 index 000000000..de4bc368e --- /dev/null +++ b/test/registered/kernels/benchmark/bench_get_block_table.py @@ -0,0 +1,49 @@ +import torch + +from sglang.kernels.jit.benchmark import marker +from sglang.kernels.jit.minicpm_sala import get_block_table +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=20, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + +_HEAD_GROUP = 2 +_SPARSE_BLOCK_SIZE = 64 +_TOPK = 96 + + +def _make_valid_inputs(token_num: int, topk: int, device: str = "cuda"): + """Well-formed inputs shared by both expansion strategies. + + ``seqlen_q_max`` is tied to ``token_num`` so the per-token causal position + (``token_pos_in_bs``) never indexes past ``block_table``. + """ + seqlen_q_max = token_num + num_blocks = max(1, seqlen_q_max // _SPARSE_BLOCK_SIZE) + torch.manual_seed(0) + topk_idx = torch.randint( + 0, num_blocks, (_HEAD_GROUP, token_num, topk), dtype=torch.int32, device=device + ) + block_table = torch.arange( + 1, seqlen_q_max + 1, dtype=torch.int32, device=device + ).reshape(1, seqlen_q_max) + token_to_bs = torch.zeros((token_num,), dtype=torch.int32, device=device) + token_pos_in_bs = torch.arange(1, token_num + 1, dtype=torch.int32, device=device) + seqlen_q = torch.tensor([seqlen_q_max], dtype=torch.int32, device=device) + return topk_idx, block_table, token_to_bs, token_pos_in_bs, seqlen_q + + +@marker.parametrize("token_num", [2**n for n in range(9, 15)], [512, 4096]) +@marker.benchmark("provider", ["blockwise", "elementwise"]) +def benchmark(token_num: int, provider: str): + inputs = _make_valid_inputs(token_num, _TOPK) + + def fn(*args): + return get_block_table(*args, elementwise=provider == "elementwise") + + return marker.do_bench(fn, input_args=inputs) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/kernels/ops/test_get_block_table.py b/test/registered/kernels/ops/test_get_block_table.py new file mode 100644 index 000000000..258463c5d --- /dev/null +++ b/test/registered/kernels/ops/test_get_block_table.py @@ -0,0 +1,153 @@ +import pytest +import torch + +from sglang.kernels.jit.minicpm_sala.get_block_table import get_block_table +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +_HEAD_GROUP = 2 +_SPARSE_BLOCK_SIZE = 64 + + +def _make_inputs(token_num, seqlen_q_max, topk, batch_size=1, device="cuda"): + """Build the same kind of inputs as the original CUDA kernel test.""" + topk_idx = torch.full( + (_HEAD_GROUP, token_num, topk), -1, dtype=torch.int32, device=device + ) + # Plant a few valid blocks at fixed positions, like the original UT. + topk_idx[0, 32, 0:2] = torch.tensor([0, 1], dtype=torch.int32, device=device) + topk_idx[1, 32, 0:2] = torch.tensor([0, 1], dtype=torch.int32, device=device) + topk_idx[1, 64, 0:2] = torch.tensor([0, 1], dtype=torch.int32, device=device) + topk_idx[0, 1000, 0:10] = torch.tensor( + [0, 1, 5, 11, 14, 16, 17, 25, 26, 27], dtype=torch.int32, device=device + ) + + block_table = torch.arange( + 1, seqlen_q_max * batch_size + 1, dtype=torch.int32, device=device + ).reshape(batch_size, seqlen_q_max) + token_to_bs = torch.zeros((token_num,), dtype=torch.int32, device=device) + token_pos_in_bs = torch.arange(1, token_num + 1, dtype=torch.int32, device=device) + seqlen_q = torch.tensor([seqlen_q_max], dtype=torch.int32, device=device) + return topk_idx, block_table, token_to_bs, token_pos_in_bs, seqlen_q + + +def _make_valid_inputs( + token_num, + seqlen_q_max, + topk, + batch_size=1, + head_group=_HEAD_GROUP, + block_size=_SPARSE_BLOCK_SIZE, + device="cuda", +): + """Build inputs with only non-negative block indices.""" + num_blocks = seqlen_q_max // block_size + torch.manual_seed(0) + topk_idx = torch.randint( + 0, num_blocks, (head_group, token_num, topk), dtype=torch.int32, device=device + ) + block_table = torch.arange( + 1, seqlen_q_max * batch_size + 1, dtype=torch.int32, device=device + ).reshape(batch_size, seqlen_q_max) + token_to_bs = torch.zeros((token_num,), dtype=torch.int32, device=device) + token_pos_in_bs = torch.arange(1, token_num + 1, dtype=torch.int32, device=device) + seqlen_q = torch.tensor([seqlen_q_max], dtype=torch.int32, device=device) + return topk_idx, block_table, token_to_bs, token_pos_in_bs, seqlen_q + + +def _get_block_table_reference( + topk_idx, + block_table, + token_to_bs, + token_pos_in_bs, + seqlen_q, + block_size=_SPARSE_BLOCK_SIZE, +): + head_group = topk_idx.shape[0] + token_num = topk_idx.shape[1] + source = topk_idx.permute(1, 0, 2).unsqueeze(-1) * block_size + torch.arange( + block_size, device=topk_idx.device + ) + valid = (source >= 0) & ( + source + < torch.minimum(seqlen_q[token_to_bs], token_pos_in_bs).view(token_num, 1, 1, 1) + ) + gathered = torch.gather( + block_table[token_to_bs], + 1, + source.reshape(token_num, -1).clamp(0, block_table.shape[1] - 1), + ).view_as(source) + heads = torch.arange(head_group, device=topk_idx.device).view(1, -1, 1, 1) + return torch.where(valid, gathered * head_group + heads, 0).flatten(2) + + +def test_get_block_table_supports_tp_local_head_group(): + inputs = _make_valid_inputs(64, 64, 96, head_group=1) + expected = _get_block_table_reference(*inputs) + actual = get_block_table(*inputs, head_group_num=1, elementwise=False) + assert torch.equal(expected, actual) + + +def _golden_check_blockwise(out_block_table, block_table, token_num): + """The assertions ported verbatim from the original kernel test.""" + # check token 32 + assert (out_block_table[32, 0] != 0).sum().item() == 33 + assert (out_block_table[32, 1] != 0).sum().item() == 33 + assert torch.equal(out_block_table[32, 0, 0:33], block_table[0][:33] * 2) + assert torch.equal(out_block_table[32, 1, 0:33], block_table[0][:33] * 2 + 1) + + # check token 64 + assert (out_block_table[64, 1] != 0).sum().item() == 65 + assert torch.equal(out_block_table[64, 1, 0:65], block_table[0][:65] * 2 + 1) + + # check token 1000 + topk_blocks = [0, 1, 5, 11, 14, 16, 17, 25, 26, 27] + tokens = [] + for b in topk_blocks: + tokens.extend(range(b * _SPARSE_BLOCK_SIZE, (b + 1) * _SPARSE_BLOCK_SIZE)) + tokens = [t for t in tokens if t < token_num and t < 1001] + assert (out_block_table[1000, 0] != 0).sum().item() == len(tokens) + assert torch.equal( + out_block_table[1000, 0, : len(tokens)], block_table[0][tokens] * 2 + ) + + +@pytest.mark.parametrize("topk", [96, 128]) +def test_get_block_table_blockwise_golden(topk): + token_num, seqlen_q_max = 8192, 8192 + inputs = _make_inputs(token_num, seqlen_q_max, topk) + out = get_block_table(*inputs, elementwise=False) + assert out.shape == (token_num, _HEAD_GROUP, topk * _SPARSE_BLOCK_SIZE) + _golden_check_blockwise(out, inputs[1], token_num) + + +@pytest.mark.parametrize("topk", [96, 128]) +def test_get_block_table_strategies_match_reference(topk): + """Both expansion strategies match the Torch reference, including -1.""" + token_num, seqlen_q_max = 2048, 2048 + inputs = _make_inputs(token_num, seqlen_q_max, topk) + expected = _get_block_table_reference(*inputs) + assert torch.equal(expected, get_block_table(*inputs, elementwise=False)) + assert torch.equal(expected, get_block_table(*inputs, elementwise=True)) + + +@pytest.mark.parametrize(("topk", "block_size"), [(10, 32), (7, 128)]) +def test_get_block_table_supports_configured_layout(topk, block_size): + token_num = seqlen_q_max = 256 + inputs = _make_valid_inputs( + token_num, + seqlen_q_max, + topk, + block_size=block_size, + ) + expected = _get_block_table_reference(*inputs, block_size=block_size) + kwargs = {"head_group_num": _HEAD_GROUP, "block_size": block_size} + assert torch.equal(expected, get_block_table(*inputs, **kwargs, elementwise=False)) + assert torch.equal(expected, get_block_table(*inputs, **kwargs, elementwise=True)) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/test_minicpm_compress_k.py b/test/registered/kernels/ops/test_minicpm_compress_k.py new file mode 100644 index 000000000..a03ea9b7d --- /dev/null +++ b/test/registered/kernels/ops/test_minicpm_compress_k.py @@ -0,0 +1,83 @@ +import pytest +import torch + +from sglang.srt.layers.attention.minicpm.fuse_kernel import ( + fused_attn_pooling_online_topk_decode, +) +from sglang.srt.layers.attention.minicpm.sparse_utils import compress_k_core_new +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +def test_compress_k_writes_each_head_once(): + """Each compressed output must be produced once even with multiple KV heads.""" + key_cache = torch.arange( + 9 * 2 * 6, + dtype=torch.float32, + device="cuda", + ).reshape(9, 2, 6) + original = key_cache.clone() + token_table = torch.tensor([[0, 0, 0, 1, 2, 3]], dtype=torch.int32, device="cuda") + compressed_table = torch.tensor([[6, 7, 8]], dtype=torch.int32, device="cuda") + full_compressed = torch.empty((3, 2, 6), device="cuda") + + compress_k_core_new( + full_compressed, + 1, + key_cache, + token_table, + compressed_table, + torch.tensor([0, 4], dtype=torch.int32, device="cuda"), + torch.tensor([1], dtype=torch.int32, device="cuda"), + torch.tensor([0, 3], dtype=torch.int32, device="cuda"), + 2, + 2, + 6, + ) + + expected = torch.stack( + ( + original[6], + original[0:2].mean(dim=0), + original[2:4].mean(dim=0), + ) + ) + torch.testing.assert_close(full_compressed, expected) + torch.testing.assert_close(key_cache[7:9], expected[1:]) + + +def test_fused_decode_topk_skips_dense_rows(): + kernel = fused_attn_pooling_online_topk_decode( + batch_size=2, + groups=16, + heads=16, + dim=128, + topk=8, + pooled_k_len=8, + dense_len=5, + dtype_str="bfloat16", + ) + topk_indices = torch.full((1, 2, 8), -1, dtype=torch.int32, device="cuda") + topk_values = torch.full( + (1, 2, 8), float("-inf"), dtype=torch.float32, device="cuda" + ) + + kernel( + torch.randn(32, 1, 128, dtype=torch.bfloat16, device="cuda"), + torch.randn(4, 1, 128, dtype=torch.bfloat16, device="cuda"), + torch.tensor([0, 1, 2], dtype=torch.int32, device="cuda"), + torch.tensor([0, 2, 4], dtype=torch.int32, device="cuda"), + torch.tensor([3, 7], dtype=torch.int32, device="cuda"), + topk_indices, + topk_values, + ) + + assert torch.all(topk_indices[:, 0] == -1) + assert torch.any(topk_indices[:, 1] >= 0) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/unit/configs/test_minicpm_config.py b/test/registered/unit/configs/test_minicpm_config.py new file mode 100644 index 000000000..9d1955bed --- /dev/null +++ b/test/registered/unit/configs/test_minicpm_config.py @@ -0,0 +1,349 @@ +import sys +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.configs.hybrid_arch import ( + hybrid_lightning_config, + mambaish_config, +) +from sglang.srt.configs.linear_attn_model_registry import ( + get_linear_attn_config, + get_linear_attn_spec_by_arch, +) +from sglang.srt.configs.mamba_utils import Mamba2CacheParams +from sglang.srt.configs.minicpm import MiniCPMHybridConfig +from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( + MambaAttnBackendBase, +) +from sglang.srt.layers.attention.linear.lightning_backend import ( + LightningAttentionBackend, +) +from sglang.srt.models import minicpm as minicpm_module +from sglang.srt.models.minicpm import ( + MiniCPMAttention, + MiniCPMDecoderLayer, + MiniCPMLightningMixer, +) +from sglang.srt.runtime_context import get_parallel +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def test_minicpm_lightning_config_defaults_are_complete(): + """A checkpoint missing optional SALA fields must still define every model input.""" + config = MiniCPMHybridConfig() + + assert config.scale_emb == 12 + assert config.scale_depth == 1.4 + assert config.dim_model_base == 256 + assert config.lightning_use_rope is True + assert config.use_output_gate is False + assert config.attention_bias is False + assert config.use_output_norm is False + assert config.qk_norm is True + + +def test_minicpm_empty_mixer_types_default_to_full_attention(): + config = MiniCPMHybridConfig(num_hidden_layers=3, mixer_types=[]) + + assert config.mixer_types == ["minicpm4", "minicpm4", "minicpm4"] + assert config.full_attention_layer_ids == [0, 1, 2] + + +def test_minicpm_sparse_config_uses_nested_fields_only(): + sparse_config = { + "block_size": 64, + "dense_len": 8192, + "init_blocks": 1, + "kernel_size": 32, + "kernel_stride": 16, + "topk": 64, + "window_size": 2048, + } + config = MiniCPMHybridConfig(sparse_config=sparse_config) + + assert config.has_minicpm_sparse_attention + assert config.sparse_config == sparse_config + assert not hasattr(config, "sparse_dense_len") + + +def test_minicpm_short_mixer_pattern_repeats_to_layer_count(): + config = MiniCPMHybridConfig( + num_hidden_layers=5, + mixer_types=["minicpm4", "lightning-attn"], + lightning_nkv=32, + ) + + assert config.mixer_types == [ + "minicpm4", + "lightning-attn", + "minicpm4", + "lightning-attn", + "minicpm4", + ] + assert config.full_attention_layer_ids == [0, 2, 4] + assert config.lightning_layer_ids == [1, 3] + + +def test_minicpm_mixer_aliases_are_canonicalized(): + config = MiniCPMHybridConfig( + num_hidden_layers=4, + mixer_types=["attention", "lightning_attn"], + lightning_nkv=32, + ) + + assert config.mixer_types == [ + "minicpm4", + "lightning-attn", + "minicpm4", + "lightning-attn", + ] + + +def test_minicpm_rejects_more_mixer_types_than_layers(): + with pytest.raises(ValueError, match="Invalid number of mixer types: 3"): + MiniCPMHybridConfig( + num_hidden_layers=2, + mixer_types=["minicpm4", "lightning", "minicpm4"], + ) + + +def test_minicpm_lightning_dimensions_fall_back_to_base_attention(): + config = MiniCPMHybridConfig( + hidden_size=96, + num_attention_heads=6, + num_key_value_heads=3, + head_dim=None, + lightning_nh=None, + lightning_nkv=None, + lightning_head_dim=None, + ) + + assert config.head_dim == 16 + assert config.lightning_nh == 6 + assert config.lightning_nkv == 3 + assert config.lightning_head_dim == 16 + + +def test_minicpm_rejects_lightning_gqa(): + with pytest.raises(ValueError, match="seg_la backend does not support GQA"): + MiniCPMHybridConfig( + num_attention_heads=6, + num_key_value_heads=3, + mixer_types=["lightning-attn"], + ) + + +def test_minicpm_lightning_idle_batch_returns_empty_output(): + """An idle DP rank must return empty output instead of reducing empty tensors.""" + mixer = MiniCPMLightningMixer.__new__(MiniCPMLightningMixer) + torch.nn.Module.__init__(mixer) + mixer.hidden_size = 8 + forward_batch = SimpleNamespace(forward_mode=SimpleNamespace(is_idle=lambda: True)) + + output = mixer.forward( + positions=torch.empty(0, dtype=torch.int64), + hidden_states=torch.empty(0, 4), + forward_batch=forward_batch, + ) + + assert output.shape == (0, 8) + + +def test_minicpm_lightning_attention_bias_applies_to_every_projection(): + """Enabling attention bias must cover every Lightning projection.""" + with get_parallel().override(tp_size=1, tp_rank=0): + mixer = MiniCPMLightningMixer( + hidden_size=8, + num_heads=2, + num_kv_heads=2, + head_dim=4, + use_rope=False, + use_output_gate=True, + attention_bias=True, + qk_norm=False, + ) + + assert mixer.qkv_proj.bias is not None + assert mixer.o_proj.bias is not None + assert mixer.z_proj.bias is not None + + +def test_minicpm_lightning_rejects_unknown_scale(): + with ( + get_parallel().override(tp_size=1, tp_rank=0), + pytest.raises(ValueError, match="Unsupported lightning scale"), + ): + MiniCPMLightningMixer( + hidden_size=8, + num_heads=2, + num_kv_heads=2, + head_dim=4, + use_rope=False, + qk_norm=False, + scale="unknown", + ) + + +def test_minicpm_full_attention_bias_applies_to_every_projection(): + """Enabling attention bias must cover every full-attention projection.""" + with get_parallel().override(tp_size=1, tp_rank=0): + mixer = MiniCPMAttention( + hidden_size=8, + num_heads=2, + num_kv_heads=2, + attn_use_rope=False, + use_output_gate=True, + attention_bias=True, + ) + + assert mixer.qkv_proj.bias is not None + assert mixer.o_proj.bias is not None + assert mixer.o_gate.bias is not None + + +def test_minicpm_full_attention_uses_configured_head_dim(monkeypatch): + monkeypatch.setattr(minicpm_module, "SiluAndMul", torch.nn.Identity) + config = MiniCPMHybridConfig( + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=2, + head_dim=6, + intermediate_size=32, + attn_use_rope=False, + ) + + with get_parallel().override(tp_size=1, tp_rank=0): + layer = MiniCPMDecoderLayer(config) + + assert layer.self_attn.head_dim == 6 + assert layer.self_attn.q_size == 12 + assert layer.self_attn.kv_size == 12 + + +def test_minicpm_lightning_reuses_shared_backend_and_cache_shape(): + config = MiniCPMHybridConfig( + num_hidden_layers=2, + mixer_types=["lightning", "minicpm4"], + lightning_nh=4, + lightning_nkv=4, + lightning_head_dim=64, + ) + + model_config = SimpleNamespace( + hf_config=config, + linear_attn_registry_result=get_linear_attn_config(config), + ) + assert hybrid_lightning_config(model_config) is config + assert mambaish_config(model_config) is config + + with get_parallel().override(attn_tp_size=1): + cache = config.mamba2_cache_params + assert isinstance(cache, Mamba2CacheParams) + assert cache.layers == [0] + assert cache.shape.conv == [(0, 0)] + assert cache.shape.temporal == (4, 64, 64) + assert config.num_linear_key_value_heads == 4 + + with get_parallel().override(attn_tp_size=1, attn_tp_rank=0): + slopes = LightningAttentionBackend._build_slope_tensor( + 4, 2, device="cpu", layerwise_decay=False + ) + assert len(slopes) == 2 + assert slopes[0].equal(slopes[1]) + + +def test_non_lightning_minicpm_is_not_classified_as_linear_attention(): + config = MiniCPMHybridConfig( + num_hidden_layers=1, + mixer_types=["minicpm4"], + sparse_config={}, + ) + model_config = SimpleNamespace( + hf_config=config, + linear_attn_registry_result=get_linear_attn_config(config), + ) + + assert hybrid_lightning_config(model_config) is None + assert mambaish_config(model_config) is None + for architecture in ("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM"): + assert get_linear_attn_spec_by_arch(architecture) is None + + +def test_lightning_backend_reads_structural_linear_config(monkeypatch): + def fake_base_init(self, model_runner): + self.topk = 1 + + monkeypatch.setattr(MambaAttnBackendBase, "__init__", fake_base_init) + config = SimpleNamespace( + num_attention_heads=8, + num_linear_key_value_heads=4, + num_hidden_layers=2, + lightning_layerwise_decay=False, + ) + model_runner = SimpleNamespace( + req_to_token_pool=SimpleNamespace( + mamba_pool=SimpleNamespace( + mamba_cache=SimpleNamespace(conv=[torch.empty(0)]) + ) + ), + sliding_window_size=None, + model_config=SimpleNamespace( + hf_config=config, + is_encoder_decoder=False, + context_len=128, + block=256, + ), + device="cpu", + kv_cache_dtype=torch.float32, + kv_cache_dtype_str="float32", + ) + + with get_parallel().override(attn_tp_size=1, attn_tp_rank=0): + backend = LightningAttentionBackend(model_runner) + + assert [slope.shape for slope in backend.tp_slope] == [(4, 1, 1), (4, 1, 1)] + assert backend.tp_slope[0].equal(backend.tp_slope[1]) + + +def test_lightning_backend_uses_layer_scale(monkeypatch): + """Each layer's attention scale must reach the linear-attention computation.""" + captured = {} + + def fake_seg_la_fwd(**kwargs): + captured.update(kwargs) + return kwargs["q"] + + monkeypatch.setattr( + "sglang.srt.layers.attention.linear.lightning_backend.seg_la_fwd", + fake_seg_la_fwd, + ) + backend = LightningAttentionBackend.__new__(LightningAttentionBackend) + backend.tp_slope = [torch.ones(1, 1, 1)] + layer = SimpleNamespace(layer_id=0, scaling=0.25) + metadata = SimpleNamespace( + batch_size=1, + query_start_loc=torch.tensor([0, 1]), + has_initial_states=torch.tensor([False]), + ) + q = torch.ones(1, 1, 1) + + backend._linear_attention_entry( + q=q, + k=q, + v=q, + kv_cache=torch.zeros(1, 1, 1, 1), + state_indices_tensor=torch.tensor([0]), + metadata=metadata, + layer=layer, + ) + + assert captured["softmax_scale"] == 0.25 + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/disaggregation/test_decode_req_to_token_pool.py b/test/registered/unit/disaggregation/test_decode_req_to_token_pool.py new file mode 100644 index 000000000..ea38e92c6 --- /dev/null +++ b/test/registered/unit/disaggregation/test_decode_req_to_token_pool.py @@ -0,0 +1,55 @@ +import sys + +import pytest +import torch + +from sglang.srt.disaggregation.decode import ( + DecodeReqToTokenPool, + HybridMambaDecodeReqToTokenPool, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +def _init_decode_pool(pool): + DecodeReqToTokenPool.__init__( + pool, + size=1, + max_context_len=4, + device="cpu", + enable_memory_saver=False, + pre_alloc_size=1, + ) + return pool + + +def test_decode_pool_reports_physical_capacity(): + pool = _init_decode_pool(DecodeReqToTokenPool.__new__(DecodeReqToTokenPool)) + + assert pool.schedulable_token_capacity(17) == 17 + + +def test_decode_pool_supports_noop_aux_cache_contract(): + pool = _init_decode_pool(DecodeReqToTokenPool.__new__(DecodeReqToTokenPool)) + req_to_token = pool.req_to_token.clone() + + pool.alloc_aux_to_lengths( + req_pool_indices_cpu=torch.tensor([1]), + target_seq_lens_cpu=torch.tensor([3]), + ) + pool.reset_aux_cache_allocator() + + assert torch.equal(pool.req_to_token, req_to_token) + + +def test_hybrid_decode_pool_initializes_aux_cache_contract(): + pool = _init_decode_pool( + HybridMambaDecodeReqToTokenPool.__new__(HybridMambaDecodeReqToTokenPool) + ) + + assert pool.schedulable_token_capacity(17) == 17 + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/layers/test_flashattention_paged_mha.py b/test/registered/unit/layers/test_flashattention_paged_mha.py new file mode 100644 index 000000000..711860095 --- /dev/null +++ b/test/registered/unit/layers/test_flashattention_paged_mha.py @@ -0,0 +1,91 @@ +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci + +with patch.dict( + sys.modules, + { + module: MagicMock() + for module in ( + "sgl_kernel", + "sgl_kernel.quantization", + "sgl_kernel.scalar_type", + ) + }, +): + from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionBackend, + ) + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _backend(): + backend = FlashAttentionBackend.__new__(FlashAttentionBackend) + backend.page_size = 1 + backend.kv_cache_dtype = torch.float16 + backend.kv_cache_dtype_str = "float8_e4m3fn" + backend.kv_cache_is_mxfp8 = False + backend.fa_impl_ver = 3 + backend.num_splits = 4 + return backend + + +class TestFlashAttentionPagedMHA(unittest.TestCase): + def test_get_paged_mha_kv_cache_supports_head_groups(self): + backend = _backend() + backend.token_to_kv_pool = SimpleNamespace( + get_kv_buffer=Mock( + return_value=( + torch.empty(8, 2, 16), + torch.empty(8, 2, 16), + ) + ) + ) + layer = SimpleNamespace( + layer_id=3, + tp_k_head_num=2, + tp_v_head_num=2, + head_dim=16, + v_head_dim=16, + ) + + key_cache, value_cache = backend.get_paged_mha_kv_cache( + layer, + head_group_num=2, + ) + + self.assertEqual(key_cache.shape, (16, 1, 1, 16)) + self.assertEqual(value_cache.shape, (16, 1, 1, 16)) + + def test_prepare_paged_mha_query_reuses_fa_scaling_policy(self): + backend = _backend() + layer = SimpleNamespace( + head_dim=16, + k_scale=torch.tensor(2.0), + v_scale=torch.tensor(4.0), + ) + q = torch.ones(2, 16, dtype=torch.bfloat16) + + q, _, _, k_descale, v_descale = backend.prepare_paged_mha_query( + q, + None, + None, + layer, + logical_batch_size=2, + kv_head_num=1, + is_prefill=True, + ) + + self.assertEqual(q.dtype, torch.float16) + self.assertEqual(k_descale.shape, (2, 1)) + self.assertEqual(v_descale.shape, (2, 1)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_minicpm_attention_adapter.py b/test/registered/unit/layers/test_minicpm_attention_adapter.py new file mode 100644 index 000000000..0c67d7ea9 --- /dev/null +++ b/test/registered/unit/layers/test_minicpm_attention_adapter.py @@ -0,0 +1,214 @@ +import sys +import unittest +from types import ModuleType, SimpleNamespace +from unittest.mock import MagicMock, Mock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci + +with patch.dict( + sys.modules, + { + module: MagicMock() + for module in ( + "sgl_kernel", + "sgl_kernel.quantization", + "sgl_kernel.scalar_type", + ) + }, +): + from sglang.srt.layers.attention.minicpm import attention_adapter as adapter_module + from sglang.srt.layers.attention.minicpm.attention_adapter import ( + MiniCPMFlashAttentionAdapter, + MiniCPMFlashInferAdapter, + ) + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _metadata(rows=1): + return SimpleNamespace( + sparse_page_table=torch.zeros((rows, 4), dtype=torch.int32), + sparse_cache_seqlens_int32=torch.full( + (rows,), + 4, + dtype=torch.int32, + ), + sparse_cu_seqlens_q=torch.arange(rows + 1, dtype=torch.int32), + sparse_cu_seqlens_k=torch.arange( + 0, + (rows + 1) * 4, + 4, + dtype=torch.int32, + ), + sparse_max_seq_len_q=1, + max_seq_len_q=1, + ) + + +class TestMiniCPMAttentionAdapter(unittest.TestCase): + def test_flashattention_adapter_owns_kernel_arguments(self): + expected = torch.ones(1, 1, 1) + flash_attn_backend = SimpleNamespace( + num_splits=4, + fa_impl_ver=3, + ) + adapter = MiniCPMFlashAttentionAdapter(flash_attn_backend) + metadata = _metadata() + layer = SimpleNamespace(scaling=0.125, logit_cap=0.0) + k_descale = torch.tensor([[2.0]]) + v_descale = torch.tensor([[4.0]]) + + with patch.object( + adapter_module, + "flash_attn_with_kvcache", + return_value=expected, + ) as kernel: + result = adapter.forward( + torch.ones(1, 1, 1), + torch.ones(4, 1, 1, 1), + torch.ones(4, 1, 1, 1), + metadata, + layer, + is_prefill=True, + k_descale=k_descale, + v_descale=v_descale, + ) + + self.assertIs(result, expected) + kwargs = kernel.call_args.kwargs + self.assertIs(kwargs["page_table"], metadata.sparse_page_table) + self.assertIs(kwargs["k_descale"], k_descale) + self.assertIs(kwargs["v_descale"], v_descale) + self.assertEqual(kwargs["num_splits"], 4) + self.assertEqual(kwargs["ver"], 3) + + def test_flashinfer_prefill_plans_once_and_executes_each_layer(self): + adapter = MiniCPMFlashInferAdapter.__new__(MiniCPMFlashInferAdapter) + adapter.prefill_planned = False + adapter._prepare = Mock() + adapter.active_rows = torch.tensor([0], dtype=torch.int32) + adapter.active_kv_indptr = torch.tensor([0, 4], dtype=torch.int32) + adapter.active_kv_indices = torch.empty(4, dtype=torch.int32) + expected = torch.ones(1, 1, 1) + adapter.active_wrapper = SimpleNamespace(forward=Mock(return_value=expected)) + metadata = _metadata() + layer = SimpleNamespace( + scaling=0.125, + logit_cap=0.0, + k_scale_float=1.0, + v_scale_float=1.0, + ) + + with patch.object( + adapter_module, + "create_flashinfer_kv_indices_triton", + ) as index_kernel: + first = adapter.forward( + torch.ones(1, 1, 1), + torch.ones(4, 1, 1, 1), + torch.ones(4, 1, 1, 1), + metadata, + layer, + is_prefill=True, + ) + second = adapter.forward( + torch.ones(1, 1, 1), + torch.ones(4, 1, 1, 1), + torch.ones(4, 1, 1, 1), + metadata, + layer, + is_prefill=True, + ) + + self.assertIs(first, expected) + self.assertIs(second, expected) + adapter._prepare.assert_called_once_with(metadata, is_prefill=True) + self.assertEqual(index_kernel.__getitem__.return_value.call_count, 2) + self.assertEqual(adapter.active_wrapper.forward.call_count, 2) + + def test_flashinfer_graph_uses_backend_wrapper_cache(self): + adapter = MiniCPMFlashInferAdapter.__new__(MiniCPMFlashInferAdapter) + adapter.device = torch.device("cpu") + adapter.head_group_num = 2 + adapter.num_qo_heads = 4 + adapter.num_kv_heads = 1 + adapter.head_dim = 16 + adapter.page_size = 1 + adapter.max_kv_tokens_per_row = 4 + adapter.q_dtype = torch.float16 + adapter.kv_dtype = torch.float16 + adapter.kv_indptr = torch.zeros(3, dtype=torch.int32) + adapter.kv_indices = torch.zeros(8, dtype=torch.int32) + adapter.kv_last_page_len = torch.ones(2, dtype=torch.int32) + adapter.rows = torch.arange(2, dtype=torch.int32) + wrapper = SimpleNamespace(begin_forward=Mock()) + adapter.flashinfer_backend = SimpleNamespace( + get_cuda_graph_decode_wrappers=Mock(return_value=[wrapper]), + ) + metadata = _metadata(rows=2) + + adapter.prepare_forward( + metadata, + is_prefill=False, + graph=True, + ) + + adapter.flashinfer_backend.get_cuda_graph_decode_wrappers.assert_called_once_with( + bs=1, + num_tokens=2, + ) + wrapper.begin_forward.assert_called_once() + self.assertIs(adapter.active_wrapper, wrapper) + + def test_flashinfer_decode_indices_cover_dense_rows(self): + wrapper = SimpleNamespace(begin_forward=Mock()) + flashinfer_backend = SimpleNamespace(decode_wrappers=[wrapper]) + flashinfer_backend_module = ModuleType( + "sglang.srt.layers.attention.flashinfer_backend" + ) + flashinfer_backend_module.FlashInferAttnBackend = Mock( + return_value=flashinfer_backend + ) + model_runner = SimpleNamespace( + device=torch.device("cpu"), + dtype=torch.float16, + kv_cache_dtype=torch.float16, + req_to_token_pool=SimpleNamespace(size=1), + ) + + with ( + patch.object(adapter_module, "is_flashinfer_available", return_value=True), + patch.dict( + sys.modules, + { + "sglang.srt.layers.attention.flashinfer_backend": ( + flashinfer_backend_module + ), + }, + ), + ): + adapter = MiniCPMFlashInferAdapter( + model_runner, + head_group_num=2, + heads_per_group=16, + head_dim=128, + page_size=1, + max_kv_tokens_per_row=7, + ) + + metadata = _metadata(rows=2) + metadata.sparse_cache_seqlens_int32.fill_(7) + adapter.prepare_forward( + metadata, + is_prefill=False, + graph=False, + ) + + self.assertEqual(adapter.kv_indices.numel(), 14) + self.assertEqual(adapter.active_kv_indices.numel(), 14) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_minicpm_sparse_cache.py b/test/registered/unit/layers/test_minicpm_sparse_cache.py new file mode 100644 index 000000000..ea471ab75 --- /dev/null +++ b/test/registered/unit/layers/test_minicpm_sparse_cache.py @@ -0,0 +1,344 @@ +import sys +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.layers.attention.minicpm.cache import ( + attach_compressed_cache, +) +from sglang.srt.managers.scheduler_components.invariant_checker import ( + SchedulerInvariantChecker, +) +from sglang.srt.managers.scheduler_components.pool_stats_observer import ( + SchedulerPoolStatsObserver, +) +from sglang.srt.mem_cache.memory_pool import ReqToTokenPool +from sglang.srt.session.streaming_session import SessionSlot, StreamingSession +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class RecordingAllocator: + def __init__(self, capacity: int): + self.capacity = capacity + self.page_size = 1 + self.next_slot = 1 + self.live: set[int] = set() + + @property + def size(self): + return self.capacity + + def alloc(self, size: int): + if size > self.available_size(): + return None + slots = torch.arange(self.next_slot, self.next_slot + size, dtype=torch.int64) + self.next_slot += size + self.live.update(slots.tolist()) + return slots + + def free(self, slots: torch.Tensor): + self.live.difference_update(slots.tolist()) + + def available_size(self): + return self.capacity - len(self.live) + + def clear(self): + self.next_slot = 1 + self.live.clear() + + +def make_pool_and_req(capacity: int = 64): + allocator = RecordingAllocator(capacity) + pool = ReqToTokenPool( + size=2, + max_context_len=64, + device="cpu", + enable_memory_saver=False, + ) + attach_compressed_cache( + pool, + allocator, + kernel_size=4, + kernel_stride=2, + enable_memory_saver=False, + ) + req = SimpleNamespace( + req_pool_idx=None, + inflight_middle_chunks=0, + kv_committed_len=0, + ) + req_pool_idx = pool.alloc([req])[0] + return pool, req, req_pool_idx, allocator + + +def alloc_extend(pool, req_pool_idx: int, seq_len: int): + pool.alloc_aux_to_lengths( + req_pool_indices_cpu=torch.tensor([req_pool_idx], dtype=torch.int64), + target_seq_lens_cpu=torch.tensor([seq_len], dtype=torch.int64), + ) + + +def test_extend_allocates_at_sparse_boundaries(): + pool, _, req_pool_idx, allocator = make_pool_and_req() + cache = pool._aux_cache + + alloc_extend(pool, req_pool_idx, seq_len=3) + assert allocator.available_size() == 39 + assert len(cache.free_slots) == 25 + + alloc_extend(pool, req_pool_idx, seq_len=4) + assert allocator.available_size() == 39 + assert len(cache.free_slots) == 24 + + alloc_extend(pool, req_pool_idx, seq_len=16) + assert allocator.available_size() == 39 + assert len(cache.free_slots) == 17 + + +def test_chunk_reuse_only_allocates_new_sparse_slots(): + pool, _, req_pool_idx, _ = make_pool_and_req() + cache = pool._aux_cache + + alloc_extend(pool, req_pool_idx, seq_len=8) + assert len(cache.free_slots) == 22 + + alloc_extend(pool, req_pool_idx, seq_len=12) + assert len(cache.free_slots) == 20 + + alloc_extend(pool, req_pool_idx, seq_len=12) + assert len(cache.free_slots) == 20 + + +def test_decode_does_not_duplicate_sparse_slots(): + """Retrying the same decode position must not allocate duplicate cache slots.""" + pool, _, req_pool_idx, _ = make_pool_and_req() + cache = pool._aux_cache + alloc_extend(pool, req_pool_idx, seq_len=15) + + pool.alloc_aux_to_lengths( + req_pool_indices_cpu=torch.tensor([req_pool_idx], dtype=torch.int64), + target_seq_lens_cpu=torch.tensor([16], dtype=torch.int64), + ) + available_after_first_decode = len(cache.free_slots) + + pool.alloc_aux_to_lengths( + req_pool_indices_cpu=torch.tensor([req_pool_idx], dtype=torch.int64), + target_seq_lens_cpu=torch.tensor([16], dtype=torch.int64), + ) + assert available_after_first_decode == 17 + assert len(cache.free_slots) == available_after_first_decode + + +def test_reserve_leaves_only_dense_capacity_visible(): + allocator = RecordingAllocator(capacity=69) + pool = ReqToTokenPool( + size=2, + max_context_len=64, + device="cpu", + enable_memory_saver=False, + ) + attach_compressed_cache( + pool, + allocator, + kernel_size=32, + kernel_stride=16, + enable_memory_saver=False, + ) + + assert allocator.available_size() == 64 + assert len(pool._aux_cache.reserved_slots) == 5 + assert pool.schedulable_token_capacity(69) == 64 + + +def test_reserved_slots_are_excluded_from_full_pool_invariant(): + pool, _, _, allocator = make_pool_and_req(capacity=69) + checker = SchedulerInvariantChecker( + is_hybrid_swa=False, + is_hybrid_ssm=True, + disaggregation_mode=None, + page_size=1, + full_tokens_per_layer=None, + swa_tokens_per_layer=None, + max_total_num_tokens=64, + tree_cache=SimpleNamespace( + supports_mamba=lambda: False, + protected_size=lambda: 0, + ), + token_to_kv_pool_allocator=allocator, + req_to_token_pool=pool, + pool_stats_observer=SimpleNamespace(session_held_tokens=lambda: 0), + get_last_batch=lambda: None, + get_running_batch=lambda: None, + ) + + leak, message = checker._check_full_pool( + SimpleNamespace( + full_available_size=allocator.available_size(), full_evictable_size=0 + ) + ) + + assert not leak, message + + +def test_hybrid_pool_stats_exclude_reserved_slots(): + pool, _, _, allocator = make_pool_and_req(capacity=69) + pool.mamba_allocator = SimpleNamespace(available_size=lambda: 1) + pool.mamba_pool = SimpleNamespace(size=1) + observer = SchedulerPoolStatsObserver( + tree_cache=SimpleNamespace(supports_mamba=lambda: False), + token_to_kv_pool_allocator=allocator, + req_to_token_pool=pool, + session_controller=None, + hisparse_coordinator=None, + is_hybrid_swa=False, + is_hybrid_ssm=True, + enable_hisparse=False, + full_tokens_per_layer=None, + swa_tokens_per_layer=None, + max_total_num_tokens=42, + get_last_batch=lambda: None, + get_running_batch=lambda: None, + ) + + stats = observer._get_mamba_token_info() + + assert stats.full_num_used == 0 + assert stats.full_token_usage == 0 + + +def test_streaming_session_release_frees_compressed_slots(): + pool, _, req_pool_idx, allocator = make_pool_and_req() + alloc_extend(pool, req_pool_idx, seq_len=16) + dense_slots = allocator.alloc(16) + pool.req_to_token[req_pool_idx, :16] = dense_slots.to(torch.int32) + compressed_cache = pool._aux_cache + assert len(compressed_cache.free_slots) < len(compressed_cache.reserved_slots) + + session = StreamingSession( + SimpleNamespace( + req_to_token_pool=pool, + token_to_kv_pool_allocator=allocator, + page_size=1, + ) + ) + session.slots["session-a"] = SessionSlot( + req_pool_idx=req_pool_idx, + kv=SimpleNamespace(kv_allocated_len=16), + ) + + session.release_session("session-a") + + assert req_pool_idx in pool.free_slots + assert len(compressed_cache.free_slots) == len(compressed_cache.reserved_slots) + + +def test_mamba_leak_diagnostic_does_not_report_reserved_slots(): + pool, _, _, allocator = make_pool_and_req(capacity=69) + allocator.free_pages = torch.arange(6, 70, dtype=torch.int64) + allocator.release_pages = torch.empty(0, dtype=torch.int64) + pool.mamba_pool = SimpleNamespace(size=1) + pool.mamba_allocator = SimpleNamespace( + size=1, + free_slots=torch.empty(0, dtype=torch.int64), + ) + checker = SchedulerInvariantChecker( + is_hybrid_swa=False, + is_hybrid_ssm=True, + disaggregation_mode=None, + page_size=1, + full_tokens_per_layer=None, + swa_tokens_per_layer=None, + max_total_num_tokens=64, + tree_cache=SimpleNamespace( + mamba_protected_size=lambda: 0, + all_values_flatten=lambda: torch.empty(0, dtype=torch.int64), + all_mamba_values_flatten=lambda: torch.empty(0, dtype=torch.int64), + ), + token_to_kv_pool_allocator=allocator, + req_to_token_pool=pool, + pool_stats_observer=SimpleNamespace( + session_held_mamba_slots=lambda: 0, + ), + get_last_batch=lambda: None, + get_running_batch=lambda: None, + ) + + leak, message = checker._check_mamba_pool( + SimpleNamespace(mamba_available_size=0, mamba_evictable_size=0) + ) + + assert leak + assert "leaked_full_pages" not in message + assert "leaked_mamba_pages={1}" in message + + +def test_partial_failure_rolls_back_and_free_releases_every_slot(): + """A failed cache-level allocation must release slots allocated for other levels.""" + pool, req, req_pool_idx, allocator = make_pool_and_req(capacity=18) + cache = pool._aux_cache + + with pytest.raises(RuntimeError, match="out of reserved slots"): + alloc_extend(pool, req_pool_idx, seq_len=16) + + assert allocator.available_size() == 11 + assert len(cache.free_slots) == 7 + + allocator.capacity = 20 + allocator.clear() + pool.reset_aux_cache_allocator() + alloc_extend(pool, req_pool_idx, seq_len=16) + assert allocator.available_size() == 12 + assert len(cache.free_slots) == 0 + + pool.free(req) + assert req.req_pool_idx is None + assert allocator.available_size() == 12 + assert len(cache.free_slots) == 8 + + +def test_allocator_reset_rebuilds_reserve(): + pool, _, _, allocator = make_pool_and_req() + + allocator.clear() + assert allocator.available_size() == 64 + + pool.reset_aux_cache_allocator() + assert allocator.available_size() == 39 + assert len(pool._aux_cache.free_slots) == 25 + + +def test_attach_compressed_cache_is_idempotent(): + allocator = RecordingAllocator(capacity=64) + pool = ReqToTokenPool( + size=2, + max_context_len=64, + device="cpu", + enable_memory_saver=False, + ) + attach_compressed_cache( + pool, + allocator, + kernel_size=4, + kernel_stride=2, + enable_memory_saver=False, + ) + cache = pool._aux_cache + k1_table = pool.req_to_sparse_k1_token + + attach_compressed_cache( + pool, + allocator, + kernel_size=4, + kernel_stride=2, + enable_memory_saver=False, + ) + + assert pool._aux_cache is cache + assert pool.req_to_sparse_k1_token is k1_table + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/layers/test_minicpm_sparse_metadata.py b/test/registered/unit/layers/test_minicpm_sparse_metadata.py new file mode 100644 index 000000000..ab0ea4b0a --- /dev/null +++ b/test/registered/unit/layers/test_minicpm_sparse_metadata.py @@ -0,0 +1,1442 @@ +import subprocess +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +with patch.dict( + sys.modules, + { + module: MagicMock() + for module in ( + "sgl_kernel", + "sgl_kernel.quantization", + "sgl_kernel.scalar_type", + ) + }, +): + from sglang.srt.layers.attention import attention_registry + from sglang.srt.layers.attention.minicpm import backend as backend_module + from sglang.srt.layers.attention.minicpm import sparse_utils + from sglang.srt.layers.attention.minicpm.attention_adapter import ( + MiniCPMFlashAttentionAdapter, + ) + from sglang.srt.layers.attention.minicpm.backend import ( + MiniCPMSparseBackend, + _gather_compressed_keys, + _transpose_head_group_layout, + ) + from sglang.srt.layers.attention.minicpm.sparse_utils import ( + CompressionLevelMetadata, + ) + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _compression_layout(): + return SimpleNamespace( + k1_kernel_size=32, + k1_kernel_stride=16, + k2_kernel_size=128, + k2_kernel_stride=64, + ) + + +def _construct_sparse_backend( + *, + max_context_len=256, + chunked_prefill_size=64, + max_running_requests=1, + use_flashinfer=False, + blackwell=False, +): + req_pool = SimpleNamespace( + req_to_sparse_k1_token=torch.empty(0), + req_to_sparse_k2_token=torch.empty(0), + ) + flash_attn_backend = SimpleNamespace( + max_context_len=max_context_len, + device="cpu", + decode_cuda_graph_metadata={}, + req_to_token_pool=req_pool, + token_to_kv_pool=SimpleNamespace(), + page_size=1, + ) + model_runner = SimpleNamespace( + dtype=torch.float16, + max_running_requests=max_running_requests, + token_to_kv_pool_allocator=SimpleNamespace(), + server_args=SimpleNamespace( + enable_memory_saver=False, + chunked_prefill_size=chunked_prefill_size, + ), + model_config=SimpleNamespace( + hf_config=SimpleNamespace( + has_minicpm_sparse_attention=True, + sparse_config={ + "kernel_size": 32, + "kernel_stride": 16, + "init_blocks": 1, + "block_size": 64, + "window_size": 64, + "dense_len": 128, + "topk": 1, + }, + ), + num_attention_heads=16, + head_dim=128, + get_num_kv_heads=lambda _tp: 1, + ), + ) + with ( + patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), + patch.object(backend_module, "is_blackwell_supported", return_value=blackwell), + patch.object( + backend_module, + "FlashAttentionBackend", + return_value=flash_attn_backend, + ) as flash_attention, + patch.object( + backend_module, + "MiniCPMFlashInferAdapter", + return_value=object(), + ), + patch.object( + backend_module, + "get_parallel", + return_value=SimpleNamespace(attn_tp_size=1), + ), + patch.object(backend_module, "attach_compressed_cache"), + ): + backend = MiniCPMSparseBackend(model_runner, use_flashinfer=use_flashinfer) + return backend, model_runner, flash_attn_backend, flash_attention + + +class _DeviceOffsetsMustNotBeRead: + def __getitem__(self, _index): + raise AssertionError("prefill layers must use scheduler-derived CPU offsets") + + +class _GraphTensorMustNotUseHostListIndex: + def __init__(self, tensor): + self.tensor = tensor + + def __getitem__(self, index): + if isinstance(index, list): + raise AssertionError("CUDA graph tensors must not use host list indices") + return self.tensor[index] + + +class _SingleTensorConversion: + def __init__(self, values): + self.values = values + self.first_item_reads = 0 + + def __len__(self): + return len(self.values) + + def __getitem__(self, index): + if index == 0: + self.first_item_reads += 1 + if self.first_item_reads > 2: + raise AssertionError("sequence lengths were converted more than once") + if index >= len(self.values): + raise IndexError + return self.values[index] + + +class TestMiniCPMSparseMetadata(CustomTestCase): + def test_sparse_backend_rejects_context_too_short_for_layout(self): + with self.assertRaisesRegex( + ValueError, + "requires context_length >= 128, got 64", + ): + _construct_sparse_backend(max_context_len=64) + + def test_fused_topk_rejects_disabled_chunked_prefill(self): + with self.assertRaisesRegex( + ValueError, + "requires a positive --chunked-prefill-size", + ): + _construct_sparse_backend( + chunked_prefill_size=-1, + use_flashinfer=True, + blackwell=True, + ) + + def test_gathered_compressed_offsets_stay_int32(self): + compressed = torch.arange(5).reshape(5, 1, 1) + level = SimpleNamespace(cu_seqlens_cpu=[0, 2, 5]) + + _, cu_seqlens = _gather_compressed_keys(compressed, level, [1]) + + self.assertEqual(cu_seqlens.dtype, torch.int32) + self.assertEqual(cu_seqlens.tolist(), [0, 3]) + + def test_registered_variants_select_adapter_explicitly(self): + runner = object() + + def build(_runner, *, use_flashinfer): + self.assertIs(_runner, runner) + return use_flashinfer + + with ( + patch.dict( + sys.modules, + { + "sglang.srt.layers.attention.minicpm.backend": backend_module, + }, + ), + patch.object( + backend_module, + "MiniCPMSparseBackend", + side_effect=build, + ), + ): + flashattn = attention_registry.ATTENTION_BACKENDS["minicpm_flashattn"]( + runner + ) + flashinfer = attention_registry.ATTENTION_BACKENDS["minicpm_flashinfer"]( + runner + ) + + self.assertFalse(flashattn) + self.assertTrue(flashinfer) + + def test_sparse_metadata_does_not_patch_base_metadata(self): + base_metadata = SimpleNamespace() + metadata_type = getattr(sparse_utils, "MiniCPMSparseMetadata") + + metadata = metadata_type(base=base_metadata) + metadata.sparse_bs_list = [0] + + self.assertEqual(metadata.sparse_bs_list, [0]) + self.assertFalse(hasattr(base_metadata, "sparse_bs_list")) + + def test_head_group_layout_round_trip(self): + tensor = torch.arange(10).reshape(5, 2, 1) + original = tensor.clone() + + _transpose_head_group_layout( + tensor, + [(1, 2)], + head_group_num=2, + heads_per_group=2, + to_group_major=True, + ) + self.assertEqual( + tensor.squeeze(-1).tolist(), + [[0, 1], [2, 3], [6, 7], [4, 5], [8, 9]], + ) + + _transpose_head_group_layout( + tensor, + [(1, 2)], + head_group_num=2, + heads_per_group=2, + to_group_major=False, + ) + self.assertTrue(torch.equal(tensor, original)) + + def test_flashattn_variant_uses_fa4_on_blackwell(self): + """Blackwell must select FA4 because FA3 binaries cannot execute there.""" + backend, model_runner, flash_attn_backend, flash_attention = ( + _construct_sparse_backend(blackwell=True) + ) + model_config = model_runner.model_config + + flash_attention.assert_called_once_with( + model_runner, + skip_prefill=False, + fa_impl_ver=4, + ) + self.assertIs(backend.flash_attn_backend, flash_attn_backend) + self.assertIs( + backend.token_to_kv_pool, + flash_attn_backend.token_to_kv_pool, + ) + self.assertIsInstance( + backend.attention_adapter, + MiniCPMFlashAttentionAdapter, + ) + self.assertEqual(backend.fused_kernel_kwargs["dtype_str"], "float16") + self.assertEqual(backend.fused_kernel_kwargs["kernel_stride"], 16) + + model_runner.server_args.attention_backend = "minicpm_flashinfer" + flashinfer_adapter = object() + with ( + patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), + patch.object(backend_module, "is_blackwell_supported", return_value=True), + patch.object( + backend_module, + "FlashAttentionBackend", + return_value=flash_attn_backend, + ), + patch.object( + backend_module, + "MiniCPMFlashInferAdapter", + return_value=flashinfer_adapter, + ), + patch.object( + backend_module, + "get_parallel", + return_value=SimpleNamespace(attn_tp_size=1), + ), + patch( + "sglang.srt.layers.attention.minicpm.fuse_kernel." + "fused_attn_pooling_online_topk_prefill", + return_value="prefill", + ), + patch.object(backend_module, "attach_compressed_cache"), + ): + backend = MiniCPMSparseBackend(model_runner, use_flashinfer=True) + + self.assertIs(backend.flash_attn_backend, flash_attn_backend) + self.assertIs(backend.attention_adapter, flashinfer_adapter) + + model_config.num_attention_heads = 8 + with ( + patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), + patch.object(backend_module, "is_blackwell_supported", return_value=True), + patch.object( + backend_module, + "FlashAttentionBackend", + return_value=flash_attn_backend, + ), + patch.object( + backend_module, + "get_parallel", + return_value=SimpleNamespace(attn_tp_size=1), + ), + patch.object(backend_module, "attach_compressed_cache"), + self.assertRaisesRegex(ValueError, "16 query heads per KV head"), + ): + MiniCPMSparseBackend(model_runner, use_flashinfer=True) + + model_runner.server_args.attention_backend = "minicpm_flashattn" + with ( + patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), + patch.object(backend_module, "is_blackwell_supported", return_value=False), + patch.object( + backend_module, + "FlashAttentionBackend", + return_value=flash_attn_backend, + ), + patch.object( + backend_module, + "get_parallel", + return_value=SimpleNamespace(attn_tp_size=1), + ), + patch.object(backend_module, "attach_compressed_cache"), + self.assertRaisesRegex(ValueError, "16 query heads per KV head"), + ): + MiniCPMSparseBackend(model_runner, use_flashinfer=False) + + def test_dense_as_sparse_routes_short_prefill(self): + req_pool = SimpleNamespace( + req_to_sparse_k1_token=torch.empty(0), + req_to_sparse_k2_token=torch.empty(0), + ) + flash_attn_backend = SimpleNamespace( + max_context_len=256, + device="cpu", + decode_cuda_graph_metadata={}, + req_to_token_pool=req_pool, + token_to_kv_pool=SimpleNamespace(), + page_size=1, + ) + hf_config = SimpleNamespace( + has_minicpm_sparse_attention=True, + sparse_config={ + "kernel_size": 32, + "kernel_stride": 16, + "init_blocks": 1, + "block_size": 64, + "window_size": 64, + "dense_len": 128, + "topk": 1, + }, + ) + model_runner = SimpleNamespace( + dtype=torch.float16, + token_to_kv_pool_allocator=SimpleNamespace(), + server_args=SimpleNamespace( + attention_backend="minicpm_flashattn", + disable_cuda_graph=False, + enable_memory_saver=False, + chunked_prefill_size=64, + ), + model_config=SimpleNamespace( + hf_config=hf_config, + num_attention_heads=16, + head_dim=128, + get_num_kv_heads=lambda _tp: 1, + ), + ) + + with ( + backend_module.envs.SGLANG_MINICPM_DENSE_AS_SPARSE.override(True), + patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), + patch.object(backend_module, "is_blackwell_supported", return_value=False), + patch.object( + backend_module, + "FlashAttentionBackend", + return_value=flash_attn_backend, + ), + patch.object( + backend_module, + "get_parallel", + return_value=SimpleNamespace(attn_tp_size=1), + ), + patch.object(backend_module, "attach_compressed_cache"), + ): + backend = MiniCPMSparseBackend(model_runner, use_flashinfer=False) + + forward_batch = SimpleNamespace( + batch_size=1, + seq_lens_cpu=torch.tensor([1], dtype=torch.int32), + seq_lens=torch.tensor([1], dtype=torch.int32), + extend_seq_lens_cpu=[1], + extend_prefix_lens_cpu=[0], + forward_mode=SimpleNamespace( + is_extend_or_draft_extend_or_mixed=lambda: True + ), + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([1], dtype=torch.int32), + page_table=torch.zeros((1, 1), dtype=torch.int32), + max_seq_len_q=1, + ) + ) + level = CompressionLevelMetadata() + with patch.object( + backend_module, + "_build_k1_k2_compression_metadata", + return_value=(level, level), + ): + backend.update_batch_for_sparse(forward_batch, metadata) + + self.assertEqual(backend.dense_len, 0) + self.assertEqual(metadata.sparse_bs_list, [0]) + + def test_dense_prefill_page_table_covers_total_sequence(self): + """Dense prefill must retain page-table coverage for the full sequence.""" + forward_batch = SimpleNamespace( + batch_size=1, + seq_lens_cpu=torch.tensor([7000], dtype=torch.int32), + extend_seq_lens_cpu=torch.tensor([2904], dtype=torch.int32), + extend_prefix_lens_cpu=[4096], + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 2904], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([7000], dtype=torch.int32), + page_table=torch.zeros((1, 7000), dtype=torch.int32), + max_seq_len_q=2904, + ) + ) + + sparse_utils._plan_sparse_prefill( + forward_batch, + metadata, + head_group_num=2, + heads_per_group=16, + dense_len=8192, + sparse_topk=96, + block_size=64, + ) + + self.assertEqual(metadata.sparse_page_table.shape, (2, 7000)) + + def test_prefill_metadata_builds_layer_invariant_cache_lengths(self): + """Sparse cache lengths must not be inferred from zero-valued table entries.""" + forward_batch = SimpleNamespace( + batch_size=2, + seq_lens_cpu=torch.tensor([200, 64], dtype=torch.int32), + extend_seq_lens_cpu=[2, 3], + extend_prefix_lens_cpu=[198, 61], + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 2, 5], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([200, 64], dtype=torch.int32), + page_table=torch.zeros((2, 200), dtype=torch.int32), + max_seq_len_q=3, + ) + ) + + sparse_utils._plan_sparse_prefill( + forward_batch, + metadata, + head_group_num=2, + heads_per_group=16, + dense_len=100, + sparse_topk=2, + block_size=64, + ) + + self.assertEqual( + metadata.sparse_cache_seqlens_int32.tolist(), + [71, 71, 72, 72, 64, 64], + ) + self.assertEqual( + metadata.sparse_cu_seqlens_k.tolist(), + [0, 71, 142, 214, 286, 350, 414], + ) + + def test_prefill_planning_builds_mixed_batch_layout(self): + forward_batch = SimpleNamespace( + batch_size=2, + seq_lens_cpu=torch.tensor([200, 64], dtype=torch.int32), + extend_seq_lens_cpu=[2, 3], + extend_prefix_lens_cpu=[198, 61], + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 2, 5], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([200, 64], dtype=torch.int32), + page_table=torch.zeros((2, 200), dtype=torch.int32), + max_seq_len_q=3, + ) + ) + + sparse_utils._plan_sparse_prefill( + forward_batch, + metadata, + head_group_num=2, + heads_per_group=16, + dense_len=100, + sparse_topk=2, + block_size=64, + ) + + self.assertEqual(metadata.sparse_bs_list, [0]) + self.assertEqual(metadata.sparse_idx, [0, 1, 2, 3]) + self.assertEqual(metadata.dense_layout, [(1, 4, 4, 3)]) + self.assertEqual(metadata.token_to_bs.tolist(), [0, 0]) + self.assertEqual(metadata.token_pos_in_bs.tolist(), [199, 200]) + self.assertEqual( + metadata.sparse_cu_seqlens_q.tolist(), + [0, 1, 2, 3, 4, 7, 10], + ) + self.assertEqual( + metadata.sparse_cache_seqlens_int32.tolist(), + [71, 71, 72, 72, 64, 64], + ) + self.assertEqual(metadata.topk_cu_seqlens_q.tolist(), [0, 2]) + self.assertEqual(metadata.topk_cu_seqlens_k.tolist(), [0, 200]) + + def test_mixed_prefill_compacts_stage1_cache_lengths(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.req_to_sparse_k1_token = torch.empty(0) + backend.req_to_sparse_k2_token = torch.empty(0) + backend.k1_kernel_size = 32 + backend.k1_kernel_stride = 16 + backend.k2_kernel_size = 128 + backend.k2_kernel_stride = 64 + backend.dense_len = 100 + backend.head_group_num = 1 + backend.sparse_topk = 2 + backend.block_size = 64 + backend.heads_per_group = 16 + + forward_batch = SimpleNamespace( + batch_size=2, + seq_lens_cpu=torch.tensor([50, 200], dtype=torch.int32), + seq_lens=torch.tensor([50, 200], dtype=torch.int32), + extend_seq_lens_cpu=[1, 1], + extend_prefix_lens_cpu=[49, 199], + req_pool_indices=torch.tensor([0, 1], dtype=torch.int64), + forward_mode=SimpleNamespace( + is_extend_or_draft_extend_or_mixed=lambda: True + ), + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([50, 200], dtype=torch.int32), + page_table=torch.arange(1, 401, dtype=torch.int32).reshape(2, 200), + max_seq_len_q=1, + ) + ) + level = CompressionLevelMetadata() + + with patch.object( + backend_module, + "_build_k1_k2_compression_metadata", + return_value=(level, level), + ): + backend.update_batch_for_sparse(forward_batch, metadata) + + self.assertEqual(metadata.sparse_bs_list, [1]) + self.assertEqual(metadata.cache_seqlens_int32_stage1.tolist(), [199]) + self.assertEqual( + metadata.sparse_page_table[0, :50].tolist(), + metadata.base.page_table[0, :50].tolist(), + ) + + def test_dense_decode_page_table_matches_batch_length(self): + forward_batch = SimpleNamespace( + batch_size=1, + seq_lens_cpu=torch.tensor([7000], dtype=torch.int32), + ) + base_metadata = SimpleNamespace( + cache_seqlens_int32=torch.tensor([7000], dtype=torch.int32), + page_table=torch.empty((1, 7000), dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + ) + + metadata = sparse_utils.MiniCPMSparseMetadata(base=base_metadata) + sparse_utils._plan_sparse_decode( + forward_batch=forward_batch, + metadata=metadata, + head_group_num=2, + dense_len=8192, + sparse_topk=96, + block_size=64, + ) + + self.assertEqual(metadata.sparse_page_table.shape, (2, 7000)) + + def test_mixed_prefill_uses_compact_sparse_page_table(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + q = torch.ones(2, 1) + k = torch.ones(2, 1, 1) + v = torch.ones(2, 1, 1) + key_cache = torch.ones(4, 1, 1, 1) + value_cache = torch.ones(4, 1, 1, 1) + backend.flash_attn_backend = SimpleNamespace( + prepare_paged_mha_query=Mock(return_value=(q, None, None, None, None)), + get_paged_mha_kv_cache=Mock(return_value=(key_cache, value_cache)), + ) + backend.token_to_kv_pool = SimpleNamespace(set_kv_buffer=Mock()) + backend.attention_adapter = SimpleNamespace( + forward=Mock(return_value=torch.ones(2, 1, 1)) + ) + backend.forward_metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + page_table=torch.tensor([[10, 0], [20, 21]], dtype=torch.int32), + ), + sparse_bs_list=[1], + sparse_idx=[1], + dense_layout=[(0, 0, 0, 1)], + sparse_page_table=torch.tensor([[10, 0], [0, 0]], dtype=torch.int32), + token_to_bs=torch.tensor([0], dtype=torch.int32), + token_pos_in_bs=torch.tensor([2], dtype=torch.int32), + seqlen_k_sparse_bs_tensor=torch.tensor([2], dtype=torch.int32), + ) + backend.head_group_num = 1 + backend.heads_per_group = 1 + backend.block_size = 1 + backend.num_sparse_topk_tokens = 1 + backend.get_topk_for_sparse = Mock( + return_value=torch.tensor([[[0]]], dtype=torch.int32) + ) + layer = SimpleNamespace( + is_cross_attention=False, + sliding_window_size=-1, + tp_q_head_num=1, + tp_k_head_num=1, + head_dim=1, + k_scale=None, + v_scale=None, + ) + forward_batch = SimpleNamespace( + batch_size=2, + seq_lens_cpu=torch.tensor([1, 2], dtype=torch.int32), + extend_seq_lens_cpu=[1, 1], + out_cache_loc=torch.tensor([0, 1], dtype=torch.int64), + forward_mode=SimpleNamespace(is_draft_extend_v2=lambda: False), + ) + + def get_sparse_page_table(_topk, page_table, *_args, **_kwargs): + self.assertEqual(page_table.tolist(), [[20, 21]]) + self.assertFalse(_kwargs["elementwise"]) + return torch.tensor([[21]], dtype=torch.int32) + + with patch.object( + backend_module, + "get_block_table", + side_effect=get_sparse_page_table, + ): + backend.forward_extend(q, k, v, layer, forward_batch) + + self.assertEqual( + backend.forward_metadata.sparse_page_table.tolist(), + [[10, 0], [21, 0]], + ) + + def test_dense_decode_copies_full_page_table(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + q = torch.ones(1, 1) + k = torch.ones(1, 1, 1) + v = torch.ones(1, 1, 1) + key_cache = torch.ones(4, 1, 1, 1) + value_cache = torch.ones(4, 1, 1, 1) + backend.flash_attn_backend = SimpleNamespace( + prepare_paged_mha_query=Mock(return_value=(q, None, None, None, None)), + get_paged_mha_kv_cache=Mock(return_value=(key_cache, value_cache)), + forward_decode=Mock(), + ) + backend.token_to_kv_pool = SimpleNamespace(set_kv_buffer=Mock()) + backend.attention_adapter = SimpleNamespace( + forward=Mock(return_value=torch.ones(1, 1, 1)) + ) + backend.forward_metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + page_table=torch.tensor([[5, 6, 7, 0]], dtype=torch.int32), + cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + max_seq_len_q=1, + ), + sparse_bs_list=[], + sparse_idx=[], + sparse_page_table=torch.tensor([[5, 6, 7, 0]], dtype=torch.int32), + sparse_cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + sparse_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + sparse_cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + token_to_bs=torch.tensor([0], dtype=torch.int32), + ) + backend.head_group_num = 1 + backend.heads_per_group = 1 + backend.page_size = 1 + backend.block_size = 1 + backend.num_sparse_topk_tokens = 2 + backend.dense_len = 4 + backend._use_cuda_graph_buffers = False + backend._compress_decode_keys = Mock() + backend.get_topk_for_sparse = Mock(return_value=None) + layer = SimpleNamespace( + is_cross_attention=False, + sliding_window_size=-1, + tp_q_head_num=1, + tp_k_head_num=1, + tp_v_head_num=1, + head_dim=1, + v_head_dim=1, + k_scale=None, + v_scale=None, + ) + forward_batch = SimpleNamespace( + batch_size=1, + seq_lens_cpu=torch.tensor([3], dtype=torch.int32), + out_cache_loc=torch.tensor([1], dtype=torch.int64), + ) + + with patch.object( + backend_module, + "get_block_table", + ) as get_block_table: + backend.forward_decode(q, k, v, layer, forward_batch) + + backend.get_topk_for_sparse.assert_called_once() + get_block_table.assert_not_called() + backend._compress_decode_keys.assert_not_called() + backend.attention_adapter.forward.assert_called_once() + backend.flash_attn_backend.forward_decode.assert_not_called() + self.assertEqual( + backend.forward_metadata.sparse_page_table[0, :3].tolist(), + [5, 6, 7], + ) + + def test_graph_decode_preserves_dense_rows_from_sparse_topk(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + q = torch.ones(2, 1) + key_cache = torch.ones(8, 1, 1, 1) + value_cache = torch.ones(8, 1, 1, 1) + backend.flash_attn_backend = SimpleNamespace( + prepare_paged_mha_query=Mock(return_value=(q, None, None, None, None)), + get_paged_mha_kv_cache=Mock(return_value=(key_cache, value_cache)), + ) + backend.token_to_kv_pool = SimpleNamespace(set_kv_buffer=Mock()) + backend.attention_adapter = SimpleNamespace( + forward=Mock(return_value=torch.ones(2, 1, 1)) + ) + backend.forward_metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + page_table=_GraphTensorMustNotUseHostListIndex( + torch.tensor([[5, 6, 7, 0, 0], [8, 9, 10, 11, 12]]) + ), + cache_seqlens_int32=torch.tensor([3, 5], dtype=torch.int32), + ), + sparse_bs_list=[0, 1], + sparse_idx=[0, 1], + token_to_bs=torch.tensor([0, 1], dtype=torch.int32), + sparse_page_table=torch.tensor([[5, 6, 7, 0, 0], [0, 0, 0, 0, 0]]), + ) + backend.head_group_num = 1 + backend.heads_per_group = 1 + backend.page_size = 1 + backend.block_size = 1 + backend.num_sparse_topk_tokens = 2 + backend.dense_len = 4 + backend._use_cuda_graph_buffers = True + backend.get_topk_for_sparse = Mock( + return_value=torch.tensor([[[0, 1], [0, 1]]], dtype=torch.int32) + ) + layer = SimpleNamespace( + is_cross_attention=False, + sliding_window_size=-1, + tp_q_head_num=1, + tp_k_head_num=1, + tp_v_head_num=1, + head_dim=1, + v_head_dim=1, + k_scale=None, + v_scale=None, + ) + forward_batch = SimpleNamespace( + batch_size=2, + out_cache_loc=torch.tensor([1, 2], dtype=torch.int64), + ) + + with patch.object( + backend_module, + "get_block_table", + return_value=torch.tensor([[30, 31], [40, 41]], dtype=torch.int32), + ): + backend.forward_decode( + q, + torch.ones(2, 1, 1), + torch.ones(2, 1, 1), + layer, + forward_batch, + ) + + self.assertEqual( + backend.forward_metadata.sparse_page_table[:, :3].tolist(), + [[5, 6, 7], [40, 41, 0]], + ) + + def test_decode_metadata_uses_scheduler_cpu_lengths(self): + """Decode metadata must not synchronize device offsets to recover lengths.""" + forward_batch = SimpleNamespace( + batch_size=2, + seq_lens_cpu=torch.tensor([64, 200], dtype=torch.int32), + ) + base_metadata = SimpleNamespace( + cache_seqlens_int32=SimpleNamespace( + dtype=torch.int32, + device=torch.device("cpu"), + ), + page_table=torch.empty((2, 200), dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + ) + + metadata = sparse_utils.MiniCPMSparseMetadata(base=base_metadata) + sparse_utils._plan_sparse_decode( + forward_batch=forward_batch, + metadata=metadata, + head_group_num=2, + dense_len=100, + sparse_topk=2, + block_size=64, + ) + + self.assertEqual( + metadata.sparse_cache_seqlens_int32.tolist(), + [64, 64, 72, 72], + ) + self.assertEqual(metadata.sparse_bs_list, [1]) + self.assertEqual(metadata.sparse_idx, [2, 3]) + self.assertEqual(metadata.dense_layout, [(0, 0, 0, 1)]) + self.assertEqual(metadata.token_to_bs.tolist(), [0]) + self.assertEqual(metadata.topk_cu_seqlens_q.tolist(), [0, 1]) + self.assertEqual(metadata.sparse_page_table.shape, (4, 128)) + + def test_cuda_graph_page_table_covers_dense_decode(self): + """Captured dense decode must reserve a threshold-sized page table.""" + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.flash_attn_backend = SimpleNamespace( + decode_cuda_graph_metadata={}, + init_cuda_graph_state=lambda *_: None, + ) + backend.attention_adapter = SimpleNamespace( + init_cuda_graph_state=lambda *_: None, + ) + backend.num_sparse_topk_tokens = 6144 + backend.page_size = 1 + backend.head_group_num = 2 + backend.device = "cpu" + backend.model_dtype = torch.float16 + backend.heads_per_group = 16 + backend.max_context_len = 256 + backend.config_dense_len = 8192 + backend.dense_len = 8192 + backend.head_dim = 128 + backend.k1_kernel_size = 32 + backend.k1_kernel_stride = 16 + backend.k2_kernel_size = 128 + backend.k2_kernel_stride = 64 + + backend.init_cuda_graph_state(max_bs=1, max_num_tokens=1) + + self.assertEqual( + backend.decode_cuda_graph_metadata["sparse_page_table"].shape, + (2, 8192), + ) + self.assertEqual( + backend.decode_cuda_graph_metadata["compress_k1"].dtype, + torch.float16, + ) + for level in ("k1", "k2"): + for field in ( + "new_token_nums", + "new_compress_token_nums", + "cu_new_compress_token_nums", + "total_compress_token_nums", + ): + self.assertNotIn( + f"{level}.{field}", + backend.decode_cuda_graph_metadata, + ) + + base_metadata = SimpleNamespace( + cu_seqlens_k=torch.zeros(2, dtype=torch.int32), + max_seq_len_k=0, + max_seq_len_q=1, + ) + capture_metadata = sparse_utils.MiniCPMSparseMetadata(base=base_metadata) + forward_batch = SimpleNamespace(batch_size=1) + backend._bind_sparse_graph_metadata( + forward_batch, + capture_metadata, + in_capture=True, + ) + self.assertIsNotNone(capture_metadata.k1) + self.assertEqual(capture_metadata.k1.cu_seqlens_cpu, [0, 511]) + self.assertEqual(capture_metadata.k2.cu_seqlens_cpu, [0, 127]) + self.assertFalse(hasattr(base_metadata, "k1")) + + base_metadata.cu_seqlens_k.copy_(torch.tensor([0, 7], dtype=torch.int32)) + replay_metadata = sparse_utils.MiniCPMSparseMetadata(base=base_metadata) + backend._bind_sparse_graph_metadata( + forward_batch, + replay_metadata, + in_capture=False, + ) + self.assertEqual(base_metadata.cu_seqlens_k.tolist(), [0, 7]) + + def test_compression_uses_configured_k1_k2_layout(self): + """K1/K2 compression must honor checkpoint strides instead of fixed defaults.""" + layer = SimpleNamespace(layer_id=0, tp_k_head_num=1, head_dim=1) + forward_batch = SimpleNamespace(req_pool_indices=[0]) + level = CompressionLevelMetadata( + table=torch.empty(0), + history_compress_token_nums=torch.empty(0), + cu_new_token_nums=torch.empty(0), + cu_total_compress_token_nums=torch.empty(0), + ) + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace(page_table=torch.empty(0)), + k1=level, + k2=level, + ) + pool = SimpleNamespace(get_key_buffer=lambda _layer_id: torch.empty(1, 1, 1)) + + with ( + patch.object(sparse_utils, "get_token_to_kv_pool", return_value=pool), + patch.object(sparse_utils, "compress_k_core_new") as compress, + ): + sparse_utils.get_compress_k_v2( + layer, + forward_batch, + metadata, + torch.empty(0), + torch.empty(0), + max_context_length=256, + k1_kernel_size=5, + k1_kernel_stride=3, + k2_kernel_size=13, + k2_kernel_stride=7, + ) + + self.assertEqual( + [(call.args[8], call.args[9]) for call in compress.call_args_list], + [(5, 3), (13, 7)], + ) + + def test_decode_compression_uses_compact_layout(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.forward_metadata = SimpleNamespace( + k1=SimpleNamespace(cu_seqlens_cpu=[0, 2, 3]), + k2=SimpleNamespace(cu_seqlens_cpu=[0, 1, 1]), + ) + backend.max_context_len = 8 + backend.k1_kernel_size = 2 + backend.k1_kernel_stride = 2 + backend.k2_kernel_size = 4 + backend.k2_kernel_stride = 4 + backend.device = torch.device("cpu") + layer = SimpleNamespace(tp_k_head_num=1, head_dim=2) + forward_batch = SimpleNamespace(batch_size=2) + + for use_graph_buffers in (False, True): + with self.subTest(use_graph_buffers=use_graph_buffers): + backend._use_cuda_graph_buffers = use_graph_buffers + backend.decode_cuda_graph_metadata = { + "compress_k1": torch.empty(8, 1, 2), + "compress_k2": torch.empty(4, 1, 2), + } + with patch.object(backend_module, "get_compress_k_v2") as compress: + k1, k2 = backend._compress_decode_keys( + torch.empty(1, dtype=torch.float16), + layer, + forward_batch, + ) + + self.assertEqual(k1.shape, (3, 1, 2)) + self.assertEqual(k2.shape, (1, 1, 2)) + compress.assert_called_once() + self.assertNotIn("padded", compress.call_args.kwargs) + + def test_dense_decode_maintains_compressed_cache_without_topk(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.forward_metadata = SimpleNamespace(sparse_bs_list=[]) + backend._compress_decode_keys = Mock( + return_value=(torch.empty(0), torch.empty(0)) + ) + backend.sparse_get_topk_impl = Mock() + layer = SimpleNamespace() + forward_batch = SimpleNamespace(batch_size=1) + + result = backend.get_topk_for_sparse( + query_states=torch.empty(1, 1, 1), + key_states=torch.empty(1, 1, 1), + layer=layer, + forward_batch=forward_batch, + is_prefill=False, + ) + + self.assertIsNone(result) + backend._compress_decode_keys.assert_called_once() + backend.sparse_get_topk_impl.assert_not_called() + + def test_mixed_decode_runs_topk_for_sparse_requests_only(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.forward_metadata = SimpleNamespace( + sparse_bs_list=[1], + base=SimpleNamespace( + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 3, 8], dtype=torch.int32), + max_seq_len_k=5, + ), + topk_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + topk_cu_seqlens_k=torch.tensor([0, 5], dtype=torch.int32), + topk_max_seqlen_k=5, + k1=SimpleNamespace( + cu_seqlens=torch.tensor([0, 1, 3], dtype=torch.int32), + cu_seqlens_cpu=[0, 1, 3], + ), + k2=SimpleNamespace( + cu_seqlens=torch.tensor([0, 1, 2], dtype=torch.int32), + cu_seqlens_cpu=[0, 1, 2], + ), + ) + backend._compress_decode_keys = Mock( + return_value=( + torch.tensor([[[10.0]], [[20.0]], [[21.0]]]), + torch.tensor([[[30.0]], [[31.0]]]), + ) + ) + backend._get_fused_topk_kernel = Mock(return_value="kernel") + backend.sparse_get_topk_impl = Mock(return_value="topk") + forward_batch = SimpleNamespace(batch_size=2) + + result = backend.get_topk_for_sparse( + query_states=torch.tensor([[[1.0]], [[2.0]]]), + key_states=torch.empty(2, 1, 1), + layer=SimpleNamespace(), + forward_batch=forward_batch, + is_prefill=False, + ) + + self.assertEqual(result, "topk") + args = backend.sparse_get_topk_impl.call_args.args + kwargs = backend.sparse_get_topk_impl.call_args.kwargs + self.assertEqual(args[0].flatten().tolist(), [2.0]) + self.assertEqual(args[1].tolist(), [0, 1]) + self.assertEqual(args[2].tolist(), [0, 3, 8]) + self.assertEqual(kwargs["compressed_k"].flatten().tolist(), [20.0, 21.0]) + self.assertEqual(kwargs["compressed_cu_seqlens"].tolist(), [0, 2]) + self.assertEqual(kwargs["compressed_k2"].flatten().tolist(), [31.0]) + self.assertEqual(kwargs["compressed_cu_seqlens2"].tolist(), [0, 1]) + backend._get_fused_topk_kernel.assert_called_once_with(1, is_prefill=False) + + def test_fused_topk_prefill_kernels_compile_for_all_batches_at_startup(self): + with patch( + "sglang.srt.layers.attention.minicpm.fuse_kernel." + "fused_attn_pooling_online_topk_prefill", + side_effect=lambda **kwargs: f"prefill-{kwargs['batch_size']}", + ): + backend, *_ = _construct_sparse_backend( + max_running_requests=3, + use_flashinfer=True, + blackwell=True, + ) + + self.assertEqual( + backend.prefill_fused_kernels, + {1: "prefill-1", 2: "prefill-2", 3: "prefill-3"}, + ) + self.assertEqual(backend.decode_fused_kernels, {}) + + def test_fused_topk_kernels_cache_each_batch_size(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.minicpm_fuse_topk = True + backend.decode_fused_kernels = {} + backend.prefill_fused_kernels = {} + backend.fused_kernel_kwargs = {"topk": 8} + backend.prefill_kernel_max_seqlen_q_grid = 64 + + with ( + patch( + "sglang.srt.layers.attention.minicpm.fuse_kernel." + "fused_attn_pooling_online_topk_prefill", + return_value="prefill", + ) as prefill, + patch( + "sglang.srt.layers.attention.minicpm.fuse_kernel." + "fused_attn_pooling_online_topk_decode", + return_value="decode", + ) as decode, + ): + self.assertEqual( + backend._get_fused_topk_kernel(3, is_prefill=True), "prefill" + ) + self.assertEqual( + backend._get_fused_topk_kernel(3, is_prefill=True), "prefill" + ) + self.assertEqual( + backend._get_fused_topk_kernel(3, is_prefill=False), "decode" + ) + self.assertEqual( + backend._get_fused_topk_kernel(3, is_prefill=False), "decode" + ) + + prefill.assert_called_once_with( + topk=8, + batch_size=3, + max_seqlen_q_grid=64, + ) + decode.assert_called_once_with(topk=8, batch_size=3) + + def test_backend_import_does_not_require_tilelang(self): + result = subprocess.run( + [ + sys.executable, + "-c", + """ +import builtins +import sys +from unittest.mock import MagicMock + +for module in ("sgl_kernel", "sgl_kernel.quantization", "sgl_kernel.scalar_type"): + sys.modules[module] = MagicMock() + +original_import = builtins.__import__ + +def import_without_tilelang(name, *args, **kwargs): + if name == "tilelang" or name.startswith("tilelang."): + raise ModuleNotFoundError(name) + return original_import(name, *args, **kwargs) + +builtins.__import__ = import_without_tilelang +import sglang.srt.layers.attention.minicpm.backend +""", + ], + capture_output=True, + text=True, + ) + self.assertEqual(result.returncode, 0, result.stderr) + + def test_forward_metadata_tracks_cuda_graph_buffer_ownership(self): + """Only replay metadata may be marked as backed by CUDA graph buffers.""" + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + metadata = SimpleNamespace() + backend.flash_attn_backend = SimpleNamespace( + forward_metadata=metadata, + init_forward_metadata=lambda *_: None, + init_forward_metadata_out_graph=lambda *_: None, + ) + backend.update_batch_for_sparse = lambda *_: None + backend._get_fused_topk_kernel = lambda *_args, **_kwargs: None + backend._bind_sparse_graph_metadata = lambda *_args, **_kwargs: None + backend._replay_sparse_graph_metadata = lambda *_: None + backend.attention_adapter = SimpleNamespace( + prepare_forward=lambda *_args, **_kwargs: None, + ) + forward_mode = SimpleNamespace( + is_target_verify=lambda: False, + is_draft_extend_v2=lambda: False, + is_idle=lambda: False, + is_decode_or_idle=lambda: True, + ) + forward_batch = SimpleNamespace(forward_mode=forward_mode, batch_size=1) + + backend._use_cuda_graph_buffers = True + backend.init_forward_metadata(forward_batch) + self.assertFalse(backend._use_cuda_graph_buffers) + self.assertIs(backend.forward_metadata.base, metadata) + + backend.init_forward_metadata_out_graph(forward_batch) + self.assertTrue(backend._use_cuda_graph_buffers) + self.assertIs(backend.forward_metadata.base, metadata) + + def test_idle_batch_skips_sparse_metadata(self): + """An idle DP rank must not attempt sparse metadata construction.""" + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + metadata = SimpleNamespace() + backend.flash_attn_backend = SimpleNamespace( + forward_metadata=metadata, + init_forward_metadata=lambda *_: None, + ) + backend.update_batch_for_sparse = lambda *_: self.fail( + "idle batches must not build sparse metadata" + ) + forward_mode = SimpleNamespace( + is_target_verify=lambda: False, + is_draft_extend_v2=lambda: False, + is_idle=lambda: True, + ) + + backend.init_forward_metadata(SimpleNamespace(forward_mode=forward_mode)) + + self.assertIs(backend.forward_metadata.base, metadata) + + def test_mixed_prefill_compiles_fused_topk_for_sparse_batch_only(self): + """A mixed batch must compile fused top-k for its sparse sub-batch only.""" + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.forward_metadata = SimpleNamespace( + sparse_bs_list=[1], + base=SimpleNamespace( + cu_seqlens_q=_DeviceOffsetsMustNotBeRead(), + ), + topk_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + topk_cu_seqlens_k=torch.tensor([0, 1], dtype=torch.int32), + topk_max_seqlen_q=1, + topk_max_seqlen_k=1, + k1=SimpleNamespace( + cu_seqlens=_DeviceOffsetsMustNotBeRead(), + cu_seqlens_cpu=[0, 0, 1], + ), + k2=SimpleNamespace( + cu_seqlens=_DeviceOffsetsMustNotBeRead(), + cu_seqlens_cpu=[0, 0, 1], + ), + ) + backend.k1_kernel_size = 1 + backend.k1_kernel_stride = 1 + backend.k2_kernel_size = 1 + backend.k2_kernel_stride = 1 + backend.dense_len = 1 + backend.max_context_len = 1 + layer = SimpleNamespace(tp_q_head_num=1, tp_k_head_num=1, head_dim=1) + forward_batch = SimpleNamespace(batch_size=2, extend_seq_lens_cpu=[1, 1]) + + with ( + patch.object( + backend_module, + "allocate_and_compress_keys", + return_value=(torch.ones(1, 1, 1), torch.ones(1, 1, 1)), + ) as allocate, + patch.object( + backend, + "_get_fused_topk_kernel", + return_value="sparse-kernel", + ) as get_kernel, + patch.object( + backend, + "sparse_get_topk_impl", + side_effect=lambda *_args, **kwargs: kwargs["fused_kernel"], + ), + ): + result = backend.get_topk_for_sparse( + query_states=torch.empty(2, 1, 1), + key_states=torch.empty(2, 1, 1), + layer=layer, + forward_batch=forward_batch, + ) + + self.assertEqual(result, "sparse-kernel") + get_kernel.assert_called_once_with(1, is_prefill=True) + + def test_compression_metadata_ignores_cuda_graph_padding(self): + """CUDA graph padding rows must not alter offsets for real requests.""" + config = _compression_layout() + + # The graph was captured for batch size 4, but only the first three + # requests are real during this replay. + forward_batch = SimpleNamespace( + batch_size=3, + seq_lens_cpu=_SingleTensorConversion([100, 200, 300]), + req_pool_indices=torch.tensor([0, 1, 2], dtype=torch.int64), + ) + base_metadata = SimpleNamespace( + cu_seqlens_q=torch.arange(5, dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 100, 200, 300, 400], dtype=torch.int32), + ) + req_to_sparse_token = torch.arange(4 * 32, dtype=torch.int32).reshape(4, 32) + + k1, k2 = sparse_utils._build_k1_k2_compression_metadata( + forward_batch=forward_batch, + base_metadata=base_metadata, + req_to_sparse_k1_token=req_to_sparse_token, + req_to_sparse_k2_token=req_to_sparse_token, + k1_kernel_size=config.k1_kernel_size, + k1_kernel_stride=config.k1_kernel_stride, + k2_kernel_size=config.k2_kernel_size, + k2_kernel_stride=config.k2_kernel_stride, + cu_seqlens_q=base_metadata.cu_seqlens_q, + ) + + self.assertEqual(k1.cu_seqlens_cpu, [0, 5, 16, 33]) + self.assertEqual(k2.cu_seqlens_cpu, [0, 0, 2, 5]) + for level in (k1, k2): + self.assertEqual(level.table.shape[0], forward_batch.batch_size) + self.assertEqual( + level.history_compress_token_nums.numel(), forward_batch.batch_size + ) + self.assertEqual(level.cu_new_token_nums.numel(), 4) + self.assertEqual(level.cu_total_compress_token_nums.numel(), 4) + + def test_sparse_graph_replay_pads_metadata_for_missing_request(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + backend.head_group_num = 2 + backend.k1_kernel_stride = 2 + backend.k2_kernel_stride = 4 + backend.max_context_len = 8 + backend.req_to_sparse_k1_token = torch.arange(8).reshape(4, 2) + backend.req_to_sparse_k2_token = torch.arange(8).reshape(4, 2) + backend.decode_cuda_graph_metadata = { + "compress_k1": torch.empty(16, 1, 1), + "compress_k2": torch.empty(8, 1, 1), + } + + def level(history, cumulative): + return CompressionLevelMetadata( + history_compress_token_nums=torch.tensor(history), + cu_seqlens=torch.tensor(cumulative), + cu_new_token_nums=torch.tensor(cumulative), + cu_total_compress_token_nums=torch.tensor(cumulative), + ) + + decode_metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace(), + sparse_idx=[2, 3, 4, 5], + dense_layout=[(0, 0, 0, 1)], + sparse_cache_seqlens_int32=torch.tensor([1, 1, 2, 2, 3, 3]), + sparse_cu_seqlens_k=torch.tensor([0, 1, 2, 4, 6, 9, 12]), + ) + compression_metadata = ( + level([1, 2, 3], [0, 1, 3, 6]), + level([4, 5, 6], [0, 4, 9, 15]), + ) + backend._build_sparse_decode_replay_metadata = Mock( + return_value=(decode_metadata, compression_metadata) + ) + + def graph_level(): + return CompressionLevelMetadata( + table=torch.full((4, 2), -1), + history_compress_token_nums=torch.full((4,), -1), + cu_seqlens=torch.full((5,), -1), + cu_new_token_nums=torch.full((5,), -1), + cu_total_compress_token_nums=torch.full((5,), -1), + ) + + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace( + cache_seqlens_int32=torch.tensor([2, 3, 4, 0]), + page_table=torch.tensor( + [[5, 6, 0, 0], [7, 8, 9, 0], [10, 11, 12, 13], [0, 0, 0, 0]] + ), + ), + k1=graph_level(), + k2=graph_level(), + sparse_page_table=torch.full((8, 4), -1, dtype=torch.int32), + sparse_cache_seqlens_int32=torch.full((8,), -1), + sparse_cu_seqlens_k=torch.full((9,), -1), + cache_seqlens_int32_stage1=torch.full((4,), -1), + ) + forward_batch = SimpleNamespace( + batch_size=4, + num_padding=1, + req_pool_indices=torch.arange(4), + seq_lens_cpu=torch.tensor([2, 3, 4, 0]), + ) + + backend._replay_sparse_graph_metadata(forward_batch, metadata) + + self.assertEqual( + metadata.sparse_cache_seqlens_int32.tolist(), [1, 1, 2, 2, 3, 3, 0, 0] + ) + self.assertEqual( + metadata.sparse_cu_seqlens_k.tolist(), [0, 1, 2, 4, 6, 9, 12, 12, 12] + ) + self.assertEqual(metadata.cache_seqlens_int32_stage1.tolist(), [1, 2, 3, 0]) + self.assertEqual(metadata.sparse_page_table[0, :2].tolist(), [10, 12]) + self.assertEqual(metadata.sparse_page_table[1, :2].tolist(), [11, 13]) + for compression, expected_history, expected_cumulative in ( + (metadata.k1, [1, 2, 3, 0], [0, 1, 3, 6, 6]), + (metadata.k2, [4, 5, 6, 0], [0, 4, 9, 15, 15]), + ): + self.assertEqual( + compression.history_compress_token_nums.tolist(), expected_history + ) + self.assertEqual(compression.cu_seqlens.tolist(), expected_cumulative) + self.assertEqual( + compression.cu_new_token_nums.tolist(), expected_cumulative + ) + self.assertEqual( + compression.cu_total_compress_token_nums.tolist(), expected_cumulative + ) + + def test_idle_sparse_graph_replay_clears_compression_lengths(self): + backend = MiniCPMSparseBackend.__new__(MiniCPMSparseBackend) + + def graph_level(): + return CompressionLevelMetadata( + history_compress_token_nums=torch.ones(2, dtype=torch.int32), + cu_seqlens=torch.ones(3, dtype=torch.int32), + cu_new_token_nums=torch.ones(3, dtype=torch.int32), + cu_total_compress_token_nums=torch.ones(3, dtype=torch.int32), + ) + + metadata = sparse_utils.MiniCPMSparseMetadata( + base=SimpleNamespace(), + k1=graph_level(), + k2=graph_level(), + sparse_cache_seqlens_int32=torch.ones(2, dtype=torch.int32), + sparse_cu_seqlens_k=torch.ones(3, dtype=torch.int32), + cache_seqlens_int32_stage1=torch.ones(2, dtype=torch.int32), + ) + + backend._replay_sparse_graph_metadata( + SimpleNamespace(batch_size=2, num_padding=2), metadata + ) + + for tensor in ( + metadata.sparse_cache_seqlens_int32, + metadata.sparse_cu_seqlens_k, + metadata.cache_seqlens_int32_stage1, + metadata.k1.history_compress_token_nums, + metadata.k1.cu_seqlens, + metadata.k1.cu_new_token_nums, + metadata.k1.cu_total_compress_token_nums, + metadata.k2.history_compress_token_nums, + metadata.k2.cu_seqlens, + metadata.k2.cu_new_token_nums, + metadata.k2.cu_total_compress_token_nums, + ): + self.assertEqual(torch.count_nonzero(tensor).item(), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_streaming_session_unit.py b/test/registered/unit/mem_cache/test_streaming_session_unit.py index 4ed28eda5..0697762e9 100644 --- a/test/registered/unit/mem_cache/test_streaming_session_unit.py +++ b/test/registered/unit/mem_cache/test_streaming_session_unit.py @@ -18,6 +18,16 @@ class _FakeAllocator: self.freed.append(free_index.clone()) +class _FakeReqToTokenPool: + def __init__(self, req_to_token): + self.req_to_token = req_to_token + self.free_slots = [] + + def free(self, req): + self.free_slots.append(req.req_pool_idx) + req.req_pool_idx = None + + class _FakeInnerCache: def __init__(self, req_to_token_pool, allocator, page_size, match_results=None): self.req_to_token_pool = req_to_token_pool @@ -85,7 +95,7 @@ def test_preabort_detaches_session_and_preserves_slot(): """Pre-aborted req (to_finish set before match_prefix) is detached from the session: session=None, abort_req() called. Slot stays intact.""" req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128) - req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[]) + req_to_token_pool = _FakeReqToTokenPool(req_to_token) allocator = _FakeAllocator() inner = _FakeInnerCache( req_to_token_pool, @@ -133,7 +143,7 @@ def test_first_mid_abort_nukes_ephemeral_slot(): slot is created from req state and nuked via release_session.""" page_size = 1 req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128) - req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[]) + req_to_token_pool = _FakeReqToTokenPool(req_to_token) allocator = _FakeAllocator() inner = _FakeInnerCache(req_to_token_pool, allocator, page_size) tree_cache = StreamingSession(inner) @@ -159,7 +169,7 @@ def test_nth_mid_abort_nukes_session_slot(): in req_nodes for next turn's re-prefill.""" page_size = 1 req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128) - req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[]) + req_to_token_pool = _FakeReqToTokenPool(req_to_token) allocator = _FakeAllocator() inner = _FakeInnerCache(req_to_token_pool, allocator, page_size) tree_cache = StreamingSession(inner) @@ -197,7 +207,7 @@ def test_release_session_threads_mamba_skip_ids(): from sglang.srt.mem_cache.unified_cache.components import ComponentType req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128) - req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[]) + req_to_token_pool = _FakeReqToTokenPool(req_to_token) allocator = _FakeAllocator() inner = _FakeInnerCache(req_to_token_pool, allocator, page_size=1) tree_cache = StreamingSession(inner) @@ -235,7 +245,7 @@ def test_trim_overshoot_postcondition(): """ page_size = 1 req_to_token = torch.arange(128, dtype=torch.int32).reshape(1, 128) - req_to_token_pool = SimpleNamespace(req_to_token=req_to_token, free_slots=[]) + req_to_token_pool = _FakeReqToTokenPool(req_to_token) allocator = _FakeAllocator() tree_cache = StreamingSession( _FakeInnerCache(req_to_token_pool, allocator, page_size) diff --git a/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py b/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py index cf3197ef5..f7a197ae0 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py +++ b/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py @@ -80,5 +80,42 @@ def test_split_full_attention_applies_model_wrapper_once(): override.restore() +def test_equal_resolved_backends_ignore_stale_global_backend(): + runner = SimpleNamespace( + server_args=SimpleNamespace( + attention_backend="global-test", + speculative_attention_mode="prefill", + ), + kv_cache_dtype=None, + token_to_kv_pool=object(), + req_to_token_pool=object(), + init_new_workspace=None, + ) + constructors = { + "global-test": lambda _runner: _FakeBackend("global"), + "resolved-test": lambda _runner: _FakeBackend("resolved"), + } + resolved = ResolvedAttentionBackendStr( + decode="resolved-test", + prefill="resolved-test", + ) + + with ( + patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors), + patch.object( + attention_backend_setup, + "attn_backend_wrapper", + side_effect=lambda _runner, backend: backend, + ), + ): + result = attention_backend_setup._build_resolved_backend( + model_runner=runner, + resolved=resolved, + init_new_workspace=False, + ) + + assert result.name == "resolved" + + if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 0c70de57c..be61c8945 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -22,6 +22,7 @@ from sglang.srt.arg_groups.overrides import ( register_model_override, validate_declarations, ) +from sglang.srt.configs.minicpm import MiniCPMHybridConfig from sglang.srt.environ import envs from sglang.srt.runtime_context import ( get_context, @@ -77,6 +78,7 @@ class TestModelOverridableWhitelist(CustomTestCase): "dcp_comm_backend", "dcp_replicate_q_proj", "disable_overlap_schedule", + "disable_radix_cache", "uses_mamba_radix_cache", "mamba_radix_cache_strategy", "mamba_full_memory_ratio", @@ -294,6 +296,235 @@ class TestGoldenModelOverrides(_IsolatedPublish): "v_head_dim": 16, } + @staticmethod + def _minicpm_overrides( + architecture, + *, + sparse_attention=False, + lightning_attention=False, + attention_backend=None, + prefill_attention_backend=None, + decode_attention_backend=None, + disaggregation_mode="null", + enable_dp_attention=False, + enable_hierarchical_cache=False, + ): + args = SimpleNamespace( + attention_backend=attention_backend, + prefill_attention_backend=prefill_attention_backend, + decode_attention_backend=decode_attention_backend, + disaggregation_mode=disaggregation_mode, + enable_dp_attention=enable_dp_attention, + enable_hierarchical_cache=enable_hierarchical_cache, + ) + args.is_attention_backend_not_set = lambda: all( + backend is None + for backend in ( + args.attention_backend, + args.prefill_attention_backend, + args.decode_attention_backend, + ) + ) + mixer_types = [] + if sparse_attention: + mixer_types.append("minicpm4") + if lightning_attention: + mixer_types.append("lightning-attn") + if not mixer_types: + mixer_types.append("minicpm4") + declarations = collect_model_override_declarations( + architecture, + args, + hf_config=MiniCPMHybridConfig( + num_hidden_layers=len(mixer_types), + num_attention_heads=1, + num_key_value_heads=1, + mixer_types=mixer_types, + sparse_config={} if sparse_attention else None, + ), + ) + return { + field: value + for _, declaration in declarations + for field, value in declaration.items() + } + + def test_minicpm_disables_radix_cache_only_for_hybrid_layers(self): + for architecture in ("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM"): + with self.subTest(architecture=architecture): + self.assertNotIn( + "disable_radix_cache", + self._minicpm_overrides(architecture), + ) + self.assertTrue( + self._minicpm_overrides(architecture, sparse_attention=True)[ + "disable_radix_cache" + ] + ) + self.assertTrue( + self._minicpm_overrides(architecture, lightning_attention=True)[ + "disable_radix_cache" + ] + ) + + def test_minicpm_rejects_dp_attention(self): + for architecture in ("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM"): + with self.subTest(architecture=architecture): + with self.assertRaisesRegex( + ValueError, + "MiniCPM does not support DP attention", + ): + self._minicpm_overrides( + architecture, + enable_dp_attention=True, + ) + + def test_minicpm_rejects_hierarchical_cache_for_hybrid_models(self): + for capability in ("sparse_attention", "lightning_attention"): + with self.subTest(capability=capability): + with self.assertRaisesRegex( + ValueError, + "MiniCPM SALA does not support hierarchical cache", + ): + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + enable_hierarchical_cache=True, + **{capability: True}, + ) + + def test_sparse_minicpm_defaults_to_sparse_attention_backend(self): + with patch.object( + overrides_module, + "is_blackwell_supported", + return_value=False, + ): + for architecture in ("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM"): + with self.subTest(architecture=architecture): + self.assertEqual( + self._minicpm_overrides( + architecture, + sparse_attention=True, + )["attention_backend"], + "minicpm_flashattn", + ) + + def test_minicpm_overrides_use_config_capabilities(self): + args = SimpleNamespace( + attention_backend=None, + prefill_attention_backend=None, + decode_attention_backend=None, + disaggregation_mode="null", + enable_dp_attention=False, + enable_hierarchical_cache=False, + is_attention_backend_not_set=lambda: True, + ) + config = SimpleNamespace( + has_minicpm_sparse_attention=True, + has_lightning_layers=False, + ) + + with patch.object( + overrides_module, "is_blackwell_supported", return_value=False + ): + overrides = overrides_module._minicpm_sala_overrides(args, config) + + self.assertTrue(overrides["disable_radix_cache"]) + self.assertEqual(overrides["attention_backend"], "minicpm_flashattn") + + def test_sparse_minicpm_defaults_to_flashinfer_on_blackwell(self): + with patch.object( + overrides_module, + "is_blackwell_supported", + return_value=True, + ): + self.assertEqual( + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + )["attention_backend"], + "minicpm_flashinfer", + ) + + def test_minicpm_preserves_explicit_attention_backend(self): + overrides = self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + attention_backend="fa3", + ) + self.assertNotIn("attention_backend", overrides) + + def test_sparse_minicpm_rejects_pd_disaggregation(self): + for disaggregation_mode in ("prefill", "decode"): + with self.subTest(disaggregation_mode=disaggregation_mode): + with self.assertRaisesRegex( + ValueError, + "MiniCPM sparse attention does not support PD disaggregation", + ): + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + disaggregation_mode=disaggregation_mode, + ) + for backend_field in ( + "prefill_attention_backend", + "decode_attention_backend", + ): + with self.subTest(backend_field=backend_field): + with self.assertRaisesRegex( + ValueError, + "MiniCPM sparse attention does not support PD disaggregation", + ): + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + disaggregation_mode="decode", + **{backend_field: "minicpm_flashattn"}, + ) + + def test_minicpm_force_dense_uses_stock_attention_backend(self): + with envs.SGLANG_MINICPM_FORCE_DENSE.override(True): + self.assertNotIn( + "attention_backend", + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + ), + ) + self.assertEqual( + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + attention_backend="minicpm_flashinfer", + )["attention_backend"], + "flashinfer", + ) + with patch.object( + overrides_module, + "is_blackwell_supported", + return_value=True, + ): + self.assertEqual( + self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + attention_backend="minicpm_flashattn", + )["attention_backend"], + "fa4", + ) + with patch.object( + overrides_module, + "is_blackwell_supported", + return_value=False, + ): + split_overrides = self._minicpm_overrides( + "MiniCPMSALAForCausalLM", + sparse_attention=True, + prefill_attention_backend="minicpm_flashattn", + decode_attention_backend="minicpm_flashattn", + ) + self.assertEqual(split_overrides["prefill_attention_backend"], "fa3") + self.assertEqual(split_overrides["decode_attention_backend"], "fa3") + def _construct(self, arch, model_type, config_extra=None, **server_kwargs): from sglang.srt.server_args import ServerArgs