[Kimi-Linear] Refactor Kimi-Linear to support RadixLinearAttention (#17506)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-01-24 21:27:13 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent bf19d20d89
commit 0c8165ffbd
3 changed files with 94 additions and 89 deletions
+41 -36
View File
@@ -16,6 +16,7 @@ from sglang.srt.distributed import (
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.attention.fla.kda import FusedRMSNormGated, fused_kda_gate
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
ColumnParallelLinear,
@@ -27,6 +28,7 @@ from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
from sglang.srt.layers.utils import PPMissingLayer
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
@@ -171,10 +173,15 @@ class KimiDeltaAttention(nn.Module):
) -> None:
super().__init__()
self.tp_size = get_tensor_model_parallel_world_size()
self.attn_tp_size = get_attention_tp_size()
self.hidden_size = hidden_size
self.config = config
self.head_dim = config.linear_attn_config["head_dim"]
self.num_heads = config.linear_attn_config["num_heads"]
self.num_k_heads = config.linear_attn_config["num_heads"]
self.num_v_heads = config.linear_attn_config["num_heads"]
self.head_k_dim = config.linear_attn_config["head_dim"]
self.head_v_dim = config.v_head_dim
self.layer_idx = layer_idx
self.prefix = prefix
assert self.num_heads % self.tp_size == 0
@@ -293,6 +300,32 @@ class KimiDeltaAttention(nn.Module):
prefix=f"{prefix}.o_proj",
)
self.q_conv_weights = self.q_conv1d.weight.view(
self.q_conv1d.weight.size(0), self.q_conv1d.weight.size(2)
)
self.k_conv_weights = self.k_conv1d.weight.view(
self.k_conv1d.weight.size(0), self.k_conv1d.weight.size(2)
)
self.v_conv_weights = self.v_conv1d.weight.view(
self.v_conv1d.weight.size(0), self.v_conv1d.weight.size(2)
)
conv_weights = (self.q_conv_weights, self.k_conv_weights, self.v_conv_weights)
bias = (self.q_conv1d.bias, self.k_conv1d.bias, self.v_conv1d.bias)
self.linear_attn = RadixLinearAttention(
layer_id=self.layer_idx,
num_qk_heads=self.num_k_heads // self.attn_tp_size,
num_v_heads=self.num_v_heads // self.attn_tp_size,
head_qk_dim=self.head_k_dim,
head_v_dim=self.head_v_dim,
attention_tp_size=self.attn_tp_size,
conv_weights=conv_weights,
bias=bias,
A_log=self.A_log,
dt_bias=self.dt_bias,
)
def forward(
self,
hidden_states: torch.Tensor,
@@ -303,54 +336,26 @@ class KimiDeltaAttention(nn.Module):
q_proj_states = self.q_proj(hidden_states)[0]
k_proj_states = self.k_proj(hidden_states)[0]
v_proj_states = self.v_proj(hidden_states)[0]
q_conv_weights = self.q_conv1d.weight.view(
self.q_conv1d.weight.size(0), self.q_conv1d.weight.size(2)
)
k_conv_weights = self.k_conv1d.weight.view(
self.k_conv1d.weight.size(0), self.k_conv1d.weight.size(2)
)
v_conv_weights = self.v_conv1d.weight.view(
self.v_conv1d.weight.size(0), self.v_conv1d.weight.size(2)
)
mixed_qkv = (q_proj_states, k_proj_states, v_proj_states)
forget_gate = self.f_b_proj(self.f_a_proj(hidden_states)[0])[0]
beta = self.b_proj(hidden_states)[0].float()
# fused_kda_gate is fused to KimiLinearAttentionBackend with decode
beta = self.b_proj(hidden_states)[0].float()
if not forward_batch.forward_mode.is_decode():
forget_gate = fused_kda_gate(
forget_gate, self.A_log, self.head_dim, g_bias=self.dt_bias
)
beta = beta.sigmoid()
forget_gate = forget_gate.unsqueeze(0)
beta = beta.unsqueeze(0)
forget_gate = forget_gate.unsqueeze(0)
kwargs = {
"q_proj_states": q_proj_states,
"k_proj_states": k_proj_states,
"v_proj_states": v_proj_states,
"q_conv_weights": q_conv_weights,
"k_conv_weights": k_conv_weights,
"v_conv_weights": v_conv_weights,
"q_conv_bias": self.q_conv1d.bias,
"k_conv_bias": self.k_conv1d.bias,
"v_conv_bias": self.v_conv1d.bias,
"head_dim": self.head_dim,
"layer_id": self.layer_idx,
"beta": beta,
"gate": forget_gate,
"A_log": self.A_log,
"dt_bias": self.dt_bias,
}
core_attn_out = forward_batch.attn_backend.forward(
q=None,
k=None,
v=None,
layer=None,
forward_batch=forward_batch,
**kwargs,
core_attn_out = self.linear_attn(
forward_batch,
mixed_qkv=mixed_qkv,
a=forget_gate,
b=beta,
)
g_proj_states = self.g_b_proj(self.g_a_proj(hidden_states)[0])[0]