[Feature] Add MiniCPM-SALA support (#30360)
Co-authored-by: Alex Nails <alex.nails@radixark.ai> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Alex Nails
Claude Opus 5
parent
d251fa2453
commit
092d85eb87
@@ -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 <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
|
||||
|
||||
#include <sgl_kernel/utils.cuh> // For LaunchKernel
|
||||
|
||||
#include <dlpack/dlpack.h>
|
||||
#include <tvm/ffi/container/tensor.h>
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
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 <int kSparseTopK, int kHeadGroup, int kSparseBlockSize>
|
||||
__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<int (*)[kHeadGroup][kSparseTopK][kSparseBlockSize]>(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 <int kSparseTopK, int kHeadGroup, int kSparseBlockSize>
|
||||
__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 <int kSparseTopK, int kHeadGroup, int kSparseBlockSize>
|
||||
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<int64_t>(kSparseTopK) * kSparseBlockSize;
|
||||
|
||||
// topk_idx: [kHeadGroup, token_num, kSparseTopK]
|
||||
TensorMatcher({static_cast<int64_t>(kHeadGroup), token_num, static_cast<int64_t>(kSparseTopK)})
|
||||
.with_dtype<int32_t>()
|
||||
.with_device<kDLCUDA>(device)
|
||||
.verify(topk_idx);
|
||||
// block_table: [batch_size, seqlen_q_max]
|
||||
TensorMatcher({batch_size, seqlen_q_max}) //
|
||||
.with_dtype<int32_t>()
|
||||
.with_device<kDLCUDA>(device)
|
||||
.verify(block_table);
|
||||
// token_to_bs / token_pos_in_bs: [token_num]
|
||||
TensorMatcher({token_num}) //
|
||||
.with_dtype<int32_t>()
|
||||
.with_device<kDLCUDA>(device)
|
||||
.verify(token_to_bs)
|
||||
.verify(token_pos_in_bs);
|
||||
// seqlen_q: [batch_size]
|
||||
TensorMatcher({batch_size}) //
|
||||
.with_dtype<int32_t>()
|
||||
.with_device<kDLCUDA>(device)
|
||||
.verify(seqlen_q);
|
||||
// out: [token_num, kHeadGroup, kSparseTopK * kSparseBlockSize]
|
||||
TensorMatcher({token_num, static_cast<int64_t>(kHeadGroup), kOutLastDim})
|
||||
.with_dtype<int32_t>()
|
||||
.with_device<kDLCUDA>(device)
|
||||
.verify(out);
|
||||
}
|
||||
|
||||
template <bool kElementwise, int kSparseTopK, int kHeadGroup, int kSparseBlockSize>
|
||||
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<kDLCUDA>();
|
||||
verify_inputs<kSparseTopK, kHeadGroup, kSparseBlockSize>(
|
||||
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<int>(token_num.unwrap());
|
||||
const int s_q_max = static_cast<int>(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<int64_t>(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<kSparseTopK, kHeadGroup, kSparseBlockSize>,
|
||||
static_cast<const int*>(topk_idx.data_ptr()),
|
||||
static_cast<const int*>(block_table.data_ptr()),
|
||||
static_cast<const int*>(token_to_bs.data_ptr()),
|
||||
static_cast<const int*>(token_pos_in_bs.data_ptr()),
|
||||
static_cast<const int*>(seqlen_q.data_ptr()),
|
||||
static_cast<int*>(out.data_ptr()),
|
||||
s_q_max,
|
||||
n_token);
|
||||
} else {
|
||||
LaunchKernel(num_blocks, kThreadsPerBlock, dev)(
|
||||
get_block_table_cuda_elementwise<kSparseTopK, kHeadGroup, kSparseBlockSize>,
|
||||
static_cast<const int*>(topk_idx.data_ptr()),
|
||||
static_cast<const int*>(block_table.data_ptr()),
|
||||
static_cast<const int*>(token_to_bs.data_ptr()),
|
||||
static_cast<const int*>(token_pos_in_bs.data_ptr()),
|
||||
static_cast<const int*>(seqlen_q.data_ptr()),
|
||||
static_cast<int*>(out.data_ptr()),
|
||||
s_q_max,
|
||||
n_token);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace sglang::minicpm_sala
|
||||
@@ -0,0 +1,3 @@
|
||||
from sglang.kernels.jit.minicpm_sala.get_block_table import get_block_table
|
||||
|
||||
__all__ = ["get_block_table"]
|
||||
@@ -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<false, {args}>",
|
||||
),
|
||||
]
|
||||
if block_size == 64 and topk % 16 == 0:
|
||||
wrappers.append(
|
||||
(
|
||||
"get_block_table_elementwise",
|
||||
f"minicpm_sala::get_block_table<true, {args}>",
|
||||
)
|
||||
)
|
||||
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
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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"]
|
||||
@@ -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).
|
||||
|
||||
@@ -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 =
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user