Opt kimi_k2_thinking biased topk module (#13150)
This commit is contained in:
@@ -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)
|
@torch.compile(dynamic=True, backend=get_compiler_backend(), disable=_is_npu)
|
||||||
def biased_grouped_topk_impl(
|
def biased_grouped_topk_impl(
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -759,6 +801,21 @@ def biased_grouped_topk_gpu(
|
|||||||
routed_scaling_factor,
|
routed_scaling_factor,
|
||||||
)
|
)
|
||||||
return topk_weights, topk_ids
|
return topk_weights, topk_ids
|
||||||
|
else:
|
||||||
|
# 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:
|
else:
|
||||||
return biased_grouped_topk_impl(
|
return biased_grouped_topk_impl(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
|
|||||||
Reference in New Issue
Block a user