[NPU]Optimization of forward_npu for UnquantizedFusedMoEMethod (#13158)

This commit is contained in:
Vladimir221
2026-02-25 08:55:20 +08:00
committed by GitHub
parent d7a03c7ebf
commit f1088beb6a
@@ -543,7 +543,6 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
layer: torch.nn.Module, layer: torch.nn.Module,
dispatch_output: StandardDispatchOutput, dispatch_output: StandardDispatchOutput,
) -> CombineInput: ) -> CombineInput:
import torch_npu
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
@@ -556,35 +555,33 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
topk_ids = topk_ids.to(torch.int32) topk_ids = topk_ids.to(torch.int32)
num_experts = layer.num_experts num_experts = layer.num_experts
top_k = layer.top_k top_k = layer.top_k
row_idx_len = num_tokens * top_k
row_idx = (
torch.arange(0, row_idx_len, dtype=torch.int32, device=topk_weights.device)
.view(top_k, -1)
.permute(1, 0)
.contiguous()
)
hidden_states, expanded_row_idx, expanded_expert_idx = ( w13 = layer.w13_weight
torch_npu.npu_moe_init_routing( w2 = layer.w2_weight
x, row_idx=row_idx, expert_idx=topk_ids, active_num=num_tokens
)
)
expert_tokens = torch_npu.npu_moe_compute_expert_tokens( hidden_states, expanded_row_idx, expert_tokens, _ = (
expanded_expert_idx, num_experts torch.ops.npu.npu_moe_init_routing_v2(
x,
topk_ids,
active_num=num_tokens * top_k,
expert_num=num_experts,
expert_tokens_num_type=1,
expert_tokens_num_flag=True,
active_expert_range=[0, num_experts],
quant_mode=-1,
)
) )
expert_tokens = expert_tokens.to(torch.int64) expert_tokens = expert_tokens.to(torch.int64)
w13_bias = [layer.w13_weight_bias] if self.with_bias else None w13_bias = [layer.w13_weight_bias] if self.with_bias else None
w2_bias = [layer.w2_weight_bias] if self.with_bias else None w2_bias = [layer.w2_weight_bias] if self.with_bias else None
# gmm1: gate_up_proj # gmm1: gate_up_proj
hidden_states = torch_npu.npu_grouped_matmul( hidden_states = torch.ops.npu.npu_grouped_matmul(
x=[hidden_states], x=[hidden_states],
weight=[layer.w13_weight], weight=[layer.w13_weight],
bias=w13_bias, bias=w13_bias,
split_item=2, split_item=2,
group_list_type=0, group_list_type=1,
group_type=0, group_type=0,
group_list=expert_tokens, group_list=expert_tokens,
output_dtype=original_dtype, output_dtype=original_dtype,
@@ -596,25 +593,25 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
hidden_states = swiglu_oai(layer, hidden_states) hidden_states = swiglu_oai(layer, hidden_states)
elif self.moe_runner_config.activation == "silu": elif self.moe_runner_config.activation == "silu":
hidden_states = torch_npu.npu_swiglu(hidden_states) hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
else: else:
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
hidden_states = GeluAndMul()(hidden_states) hidden_states = GeluAndMul()(hidden_states)
# gmm2: down_proj # gmm2: down_proj
hidden_states = torch_npu.npu_grouped_matmul( hidden_states = torch.ops.npu.npu_grouped_matmul(
x=[hidden_states], x=[hidden_states],
weight=[layer.w2_weight], weight=[layer.w2_weight],
bias=w2_bias, bias=w2_bias,
split_item=2, split_item=2,
group_list_type=0, group_list_type=1,
group_type=0, group_type=0,
group_list=expert_tokens, group_list=expert_tokens,
output_dtype=original_dtype, output_dtype=original_dtype,
)[0] )[0]
final_hidden_states = torch_npu.npu_moe_finalize_routing( final_hidden_states = torch.ops.npu.npu_moe_finalize_routing(
hidden_states, hidden_states,
skip1=None, skip1=None,
skip2=None, skip2=None,
@@ -622,6 +619,7 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
scales=topk_weights, scales=topk_weights,
expanded_src_to_dst_row=expanded_row_idx, expanded_src_to_dst_row=expanded_row_idx,
export_for_source_row=topk_ids, export_for_source_row=topk_ids,
drop_pad_mode=2,
) )
return StandardCombineInput(hidden_states=final_hidden_states) return StandardCombineInput(hidden_states=final_hidden_states)