[NPU] Support dequant_swiglu_quant & moe_init_routing_v2 & npu_moe_token_unpermute for W8A8 MoE decode (#19913)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user