diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 765224008..e0ddbcd05 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -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, diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 5d93cb735..8a9029ce4 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -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]