From a1cb717d0b4a4d7c5248a7a0d3b2a26599a40f7b Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Thu, 13 Nov 2025 08:52:06 +0800 Subject: [PATCH] Opt kimi_k2_thinking biased topk module (#13150) --- python/sglang/srt/layers/moe/topk.py | 85 +++++++++++++++++++++++----- 1 file changed, 71 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 203cd5f41..006552e37 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -600,6 +600,48 @@ def grouped_topk_cpu( ) +@torch.compile(dynamic=True, backend=get_compiler_backend(), disable=_is_npu) +def kimi_k2_biased_topk_impl( + hidden_states: torch.Tensor, + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + topk: int, + renormalize: bool, + routed_scaling_factor: Optional[float] = None, + num_token_non_padded: Optional[torch.Tensor] = None, + expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None, + apply_routed_scaling_factor_on_output: Optional[bool] = False, +): + """ + Optimized version for num_expert_group=1 case (e.g., Kimi K2 with 384 experts). + Simplifies the grouped topk logic by removing unnecessary group masking operations. + Note: This function assumes num_fused_shared_experts=0. + """ + assert hidden_states.shape[0] == gating_output.shape[0], "Number of tokens mismatch" + + scores = gating_output.sigmoid() + num_token = scores.shape[0] + + # When num_expert_group=1, no need for group masking + # Directly compute scores with correction bias + tmp_scores = scores.view(num_token, -1) + correction_bias.unsqueeze(0) + + # Directly select topk experts (no need to sort since num_fused_shared_experts=0) + _, topk_ids = torch.topk(tmp_scores, k=topk, dim=-1, sorted=False) + topk_weights = scores.gather(1, topk_ids) + + if renormalize: + topk_weights_sum = topk_weights.sum(dim=-1, keepdim=True) + topk_weights = topk_weights / topk_weights_sum + if apply_routed_scaling_factor_on_output: + topk_weights *= routed_scaling_factor + + topk_weights, topk_ids = topk_weights.to(torch.float32), topk_ids.to(torch.int32) + topk_ids = topk_ids_logical_to_physical(topk_ids, expert_location_dispatch_info) + _mask_topk_ids_padded_region(topk_ids, num_token_non_padded) + return topk_weights, topk_ids + + @torch.compile(dynamic=True, backend=get_compiler_backend(), disable=_is_npu) def biased_grouped_topk_impl( hidden_states: torch.Tensor, @@ -760,20 +802,35 @@ def biased_grouped_topk_gpu( ) return topk_weights, topk_ids else: - return biased_grouped_topk_impl( - hidden_states, - gating_output, - correction_bias, - topk, - renormalize, - num_expert_group, - topk_group, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - num_token_non_padded=num_token_non_padded, - expert_location_dispatch_info=expert_location_dispatch_info, - apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output, - ) + # Use optimized path for Kimi K2 (384 experts with num_expert_group=1) + num_experts = gating_output.shape[1] + if num_experts == 384 and num_expert_group == 1: + return kimi_k2_biased_topk_impl( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + routed_scaling_factor=routed_scaling_factor, + num_token_non_padded=num_token_non_padded, + expert_location_dispatch_info=expert_location_dispatch_info, + apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output, + ) + else: + return biased_grouped_topk_impl( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group, + topk_group, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + num_token_non_padded=num_token_non_padded, + expert_location_dispatch_info=expert_location_dispatch_info, + apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output, + ) def biased_grouped_topk_cpu(