[NPU] bugfix for Qwen3-Next and performance update (#11969)
This commit is contained in:
@@ -314,16 +314,41 @@ class TopK(CustomOp):
|
||||
num_token_non_padded: Optional[torch.Tensor] = None,
|
||||
expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None,
|
||||
) -> TopKOutput:
|
||||
global_num_experts = router_logits.shape[-1]
|
||||
|
||||
# NOTE: now npu_moe_gating_top_k can only support `group_count=256` pattern
|
||||
if global_num_experts == 256:
|
||||
use_grouped_topk = self.topk_config.use_grouped_topk
|
||||
torch_native = self.topk_config.torch_native
|
||||
renormalize = self.topk_config.renormalize
|
||||
|
||||
if not use_grouped_topk and not torch_native:
|
||||
topk_weights, topk_ids, _ = torch_npu.npu_moe_gating_top_k_softmax(
|
||||
router_logits,
|
||||
k=self.topk_config.top_k,
|
||||
)
|
||||
topk_weights = topk_weights.to(torch.float32)
|
||||
|
||||
if renormalize:
|
||||
topk_weights_sum = (
|
||||
topk_weights.sum(dim=-1, keepdim=True)
|
||||
if self.topk_config.num_fused_shared_experts == 0
|
||||
else topk_weights[:, :-1].sum(dim=-1, keepdim=True)
|
||||
)
|
||||
topk_weights = topk_weights / topk_weights_sum
|
||||
|
||||
if expert_location_dispatch_info is not None:
|
||||
topk_ids = topk_ids_logical_to_physical(
|
||||
topk_ids, expert_location_dispatch_info
|
||||
)
|
||||
get_global_expert_distribution_recorder().on_select_experts(
|
||||
topk_ids=topk_ids
|
||||
)
|
||||
|
||||
return StandardTopKOutput(topk_weights, topk_ids, _)
|
||||
if use_grouped_topk and not torch_native and router_logits.shape[-1] == 256:
|
||||
# NOTE: now npu_moe_gating_top_k can only support `group_count=256` pattern
|
||||
routed_scaling_factor = self.topk_config.routed_scaling_factor or 1
|
||||
router_logits = router_logits.to(torch.float32)
|
||||
|
||||
topk_weights, topk_ids, _ = torch_npu.npu_moe_gating_top_k(
|
||||
router_logits,
|
||||
router_logits.to(torch.float32),
|
||||
k=self.topk_config.top_k,
|
||||
bias=self.topk_config.correction_bias.to(torch.float32),
|
||||
k_group=self.topk_config.topk_group,
|
||||
@@ -335,7 +360,7 @@ class TopK(CustomOp):
|
||||
eps=float(1e-20),
|
||||
)
|
||||
|
||||
if self.topk_config.renormalize:
|
||||
if renormalize:
|
||||
topk_weights_sum = (
|
||||
topk_weights.sum(dim=-1, keepdim=True)
|
||||
if self.topk_config.num_fused_shared_experts == 0
|
||||
|
||||
Reference in New Issue
Block a user