[Kimi-Linear] Refactor kimi-linear gate calculation to avoid duplicated code (#17160)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-01-20 14:29:24 +08:00
committed by GitHub
parent a3addd6203
commit e6b7c04947
2 changed files with 18 additions and 42 deletions

View File

@@ -17,11 +17,7 @@ from sglang.srt.layers.attention.fla.fused_recurrent import (
from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import (
fused_sigmoid_gating_delta_rule_update,
)
from sglang.srt.layers.attention.fla.kda import (
chunk_kda,
fused_kda_gate,
fused_recurrent_kda,
)
from sglang.srt.layers.attention.fla.kda import chunk_kda, fused_recurrent_kda
from sglang.srt.layers.attention.mamba.causal_conv1d_triton import (
PAD_SLOT_ID,
causal_conv1d_fn,
@@ -646,14 +642,10 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
k_conv_bias = kwargs["k_conv_bias"]
v_conv_bias = kwargs["v_conv_bias"]
A_log = kwargs["A_log"]
dt_bias = kwargs["dt_bias"]
b_proj = kwargs["b_proj"]
f_a_proj = kwargs["f_a_proj"]
f_b_proj = kwargs["f_b_proj"]
hidden_states = kwargs["hidden_states"]
head_dim = kwargs["head_dim"]
layer_id = kwargs["layer_id"]
beta = kwargs["beta"]
g = kwargs["gate"]
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
q_conv_state, k_conv_state, v_conv_state = layer_cache.conv
@@ -694,14 +686,6 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
lambda x: rearrange(x, "n (h d) -> 1 n h d", d=head_dim), (q, k, v)
)
beta = b_proj(hidden_states)[0].float().sigmoid()
g = f_b_proj(f_a_proj(hidden_states)[0])[0]
g = fused_kda_gate(g, A_log, head_dim, g_bias=dt_bias)
beta = beta.unsqueeze(0)
g = g.unsqueeze(0)
initial_state = ssm_states[cache_indices].contiguous()
(
core_attn_out,
@@ -744,14 +728,10 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
k_conv_bias = kwargs["k_conv_bias"]
v_conv_bias = kwargs["v_conv_bias"]
A_log = kwargs["A_log"]
dt_bias = kwargs["dt_bias"]
b_proj = kwargs["b_proj"]
f_a_proj = kwargs["f_a_proj"]
f_b_proj = kwargs["f_b_proj"]
hidden_states = kwargs["hidden_states"]
head_dim = kwargs["head_dim"]
layer_id = kwargs["layer_id"]
beta = kwargs["beta"]
g = kwargs["gate"]
query_start_loc = self.forward_metadata.query_start_loc
cache_indices = self.forward_metadata.mamba_cache_indices
@@ -811,14 +791,6 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
lambda x: rearrange(x, "n (h d) -> 1 n h d", d=head_dim), (q, k, v)
)
beta = b_proj(hidden_states)[0].float().sigmoid()
g = f_b_proj(f_a_proj(hidden_states)[0])[0]
g = fused_kda_gate(g, A_log, head_dim, g_bias=dt_bias)
beta = beta.unsqueeze(0)
g = g.unsqueeze(0)
core_attn_out = chunk_kda(
q=q,
k=k,

View File

@@ -15,7 +15,7 @@ from sglang.srt.distributed import (
tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.attention.fla.kda import FusedRMSNormGated
from sglang.srt.layers.attention.fla.kda import FusedRMSNormGated, fused_kda_gate
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
ColumnParallelLinear,
@@ -314,6 +314,14 @@ class KimiDeltaAttention(nn.Module):
self.v_conv1d.weight.size(0), self.v_conv1d.weight.size(2)
)
beta = self.b_proj(hidden_states)[0].float().sigmoid()
forget_gate = self.f_b_proj(self.f_a_proj(hidden_states)[0])[0]
forget_gate = fused_kda_gate(
forget_gate, self.A_log, self.head_dim, g_bias=self.dt_bias
)
beta = beta.unsqueeze(0)
forget_gate = forget_gate.unsqueeze(0)
kwargs = {
"q_proj_states": q_proj_states,
"k_proj_states": k_proj_states,
@@ -324,14 +332,10 @@ class KimiDeltaAttention(nn.Module):
"q_conv_bias": self.q_conv1d.bias,
"k_conv_bias": self.k_conv1d.bias,
"v_conv_bias": self.v_conv1d.bias,
"dt_bias": self.dt_bias,
"b_proj": self.b_proj,
"f_a_proj": self.f_a_proj,
"f_b_proj": self.f_b_proj,
"A_log": self.A_log,
"head_dim": self.head_dim,
"hidden_states": hidden_states,
"layer_id": self.layer_idx,
"beta": beta,
"gate": forget_gate,
}
core_attn_out = forward_batch.attn_backend.forward(
@@ -344,8 +348,8 @@ class KimiDeltaAttention(nn.Module):
)
g_proj_states = self.g_b_proj(self.g_a_proj(hidden_states)[0])[0]
g = rearrange(g_proj_states, "... (h d) -> ... h d", d=self.head_dim)
core_attn_out = self.o_norm(core_attn_out, g)
norm_gate = rearrange(g_proj_states, "... (h d) -> ... h d", d=self.head_dim)
core_attn_out = self.o_norm(core_attn_out, norm_gate)
core_attn_out = rearrange(core_attn_out, "1 n h d -> n (h d)")
return self.o_proj(core_attn_out)[0]