[NPU] Support dequant_swiglu_quant & moe_init_routing_v2 & npu_moe_token_unpermute for W8A8 MoE decode (#19913)

This commit is contained in:
heziiop
2026-03-17 21:39:29 +08:00
committed by GitHub
parent 5717834f1f
commit b5f3eaecbc

View File

@@ -162,15 +162,19 @@ def npu_fused_experts(
output_dtype=original_dtype,
)[0]
# act_fn: swiglu
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
if not use_wna16:
hidden_states, pertoken_scale = torch.ops.npu.npu_dynamic_quant(hidden_states)
hidden_states, pertoken_scale = torch.ops.npu.npu_dequant_swiglu_quant(
hidden_states,
activate_left=True,
quant_mode=1,
)
scale_args2 = {
"scale": [w2_scale.to(scale_dtype)],
"per_token_scale": [pertoken_scale],
}
else:
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
scale_args2 = {"antiquant_scale": [w2_scale], "antiquant_offset": [w2_offset]}
# gmm2: down_proj
hidden_states = torch.ops.npu.npu_grouped_matmul(
@@ -198,6 +202,78 @@ def npu_fused_experts(
return final_hidden_states
def npu_fused_experts_w8a8_decode(
hidden_states: torch.Tensor,
w13: torch.Tensor,
w13_scale: torch.Tensor,
w2: torch.Tensor,
w2_scale: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
top_k: int,
**kwargs,
):
num_tokens = hidden_states.shape[:-1].numel()
first_expert_idx = 0
last_expert_idx = w13.shape[0]
global_num_experts = w13.shape[0]
original_shape = hidden_states.shape
group_list_type = 1
sorted_hidden_states, expanded_row_idx, expert_tokens, pertoken_scale = (
torch.ops.npu.npu_moe_init_routing_v2(
hidden_states,
topk_ids,
active_num=num_tokens * top_k,
expert_num=global_num_experts,
expert_tokens_num_type=group_list_type,
expert_tokens_num_flag=True,
active_expert_range=[first_expert_idx, last_expert_idx],
quant_mode=1,
)
)
hidden_states = torch.ops.npu.npu_grouped_matmul(
x=[sorted_hidden_states],
weight=[w13],
scale=[w13_scale],
per_token_scale=[pertoken_scale],
group_list=expert_tokens,
split_item=2,
group_type=0,
group_list_type=group_list_type,
output_dtype=torch.bfloat16,
)[0]
# act_fn: swiglu
hidden_states, swiglu_out_scale = torch.ops.npu.npu_dequant_swiglu_quant(
hidden_states, quant_mode=1, activate_left=True
)
output = torch.ops.npu.npu_grouped_matmul(
x=[hidden_states],
weight=[w2],
scale=[w2_scale],
per_token_scale=[swiglu_out_scale],
group_list=expert_tokens,
split_item=2,
group_type=0,
group_list_type=group_list_type,
output_dtype=torch.bfloat16,
)[0]
assert original_shape is not None
final_hidden_states = torch.ops.npu.npu_moe_token_unpermute(
permuted_tokens=output,
sorted_indices=torch.abs(expanded_row_idx),
probs=topk_weights,
)
if len(original_shape) == 3:
final_hidden_states = final_hidden_states.view(original_shape)
return final_hidden_states
def npu_fused_moe_without_routing_weights_bf16(
layer, hidden_states, group_list_type, group_list, output_dtype
):
@@ -396,6 +472,12 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
layer.w2_weight_scale = torch.nn.Parameter(
layer.w2_weight_scale.data.squeeze(-1), requires_grad=False
)
layer.w13_weight_scale_bf16 = torch.nn.Parameter(
layer.w13_weight_scale.data.to(dtype=torch.bfloat16), requires_grad=False
)
layer.w2_weight_scale_bf16 = torch.nn.Parameter(
layer.w2_weight_scale.data.to(dtype=torch.bfloat16), requires_grad=False
)
# Compressed-tensors format doesn't have this field
if hasattr(layer, "w13_weight_offset"):
layer.w13_weight_offset = torch.nn.Parameter(
@@ -415,22 +497,42 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
) -> "CombineInput":
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
x = dispatch_output.hidden_states
# release fp32 scale to save memory
layer.w13_weight_scale = None
layer.w2_weight_scale = None
hidden_states = dispatch_output.hidden_states
topk_output = dispatch_output.topk_output
topk_weights, topk_ids, _ = topk_output
topk_ids = topk_ids.to(torch.int32)
topk_weights = topk_weights.to(x.dtype)
output = npu_fused_experts(
hidden_states=x,
w13=layer.w13_weight,
w13_scale=layer.w13_weight_scale,
w2=layer.w2_weight,
w2_scale=layer.w2_weight_scale,
topk_weights=topk_weights,
topk_ids=topk_ids,
top_k=topk_ids.shape[1],
)
topk_weights = topk_weights.to(hidden_states.dtype)
# prefill
if not torch.npu.is_current_stream_capturing():
output = npu_fused_experts(
hidden_states=hidden_states,
w13=layer.w13_weight,
w13_scale=layer.w13_weight_scale_bf16,
w2=layer.w2_weight,
w2_scale=layer.w2_weight_scale_bf16,
topk_weights=topk_weights,
topk_ids=topk_ids,
top_k=topk_ids.shape[1],
)
# decode
else:
output = npu_fused_experts_w8a8_decode(
hidden_states=hidden_states,
w13=layer.w13_weight,
w13_scale=layer.w13_weight_scale_bf16,
w2=layer.w2_weight,
w2_scale=layer.w2_weight_scale_bf16,
topk_weights=topk_weights,
topk_ids=topk_ids,
top_k=topk_ids.shape[1],
)
return StandardCombineInput(hidden_states=output)
def apply_without_routing_weights(