[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,
|
output_dtype=original_dtype,
|
||||||
)[0]
|
)[0]
|
||||||
# act_fn: swiglu
|
# act_fn: swiglu
|
||||||
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
|
|
||||||
if not use_wna16:
|
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_args2 = {
|
||||||
"scale": [w2_scale.to(scale_dtype)],
|
"scale": [w2_scale.to(scale_dtype)],
|
||||||
"per_token_scale": [pertoken_scale],
|
"per_token_scale": [pertoken_scale],
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
|
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
|
||||||
scale_args2 = {"antiquant_scale": [w2_scale], "antiquant_offset": [w2_offset]}
|
scale_args2 = {"antiquant_scale": [w2_scale], "antiquant_offset": [w2_offset]}
|
||||||
# gmm2: down_proj
|
# gmm2: down_proj
|
||||||
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
||||||
@@ -198,6 +202,78 @@ def npu_fused_experts(
|
|||||||
return final_hidden_states
|
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(
|
def npu_fused_moe_without_routing_weights_bf16(
|
||||||
layer, hidden_states, group_list_type, group_list, output_dtype
|
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 = torch.nn.Parameter(
|
||||||
layer.w2_weight_scale.data.squeeze(-1), requires_grad=False
|
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
|
# Compressed-tensors format doesn't have this field
|
||||||
if hasattr(layer, "w13_weight_offset"):
|
if hasattr(layer, "w13_weight_offset"):
|
||||||
layer.w13_weight_offset = torch.nn.Parameter(
|
layer.w13_weight_offset = torch.nn.Parameter(
|
||||||
@@ -415,22 +497,42 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
) -> "CombineInput":
|
) -> "CombineInput":
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
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_output = dispatch_output.topk_output
|
||||||
|
|
||||||
topk_weights, topk_ids, _ = topk_output
|
topk_weights, topk_ids, _ = topk_output
|
||||||
topk_ids = topk_ids.to(torch.int32)
|
topk_ids = topk_ids.to(torch.int32)
|
||||||
topk_weights = topk_weights.to(x.dtype)
|
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||||
output = npu_fused_experts(
|
|
||||||
hidden_states=x,
|
# prefill
|
||||||
w13=layer.w13_weight,
|
if not torch.npu.is_current_stream_capturing():
|
||||||
w13_scale=layer.w13_weight_scale,
|
output = npu_fused_experts(
|
||||||
w2=layer.w2_weight,
|
hidden_states=hidden_states,
|
||||||
w2_scale=layer.w2_weight_scale,
|
w13=layer.w13_weight,
|
||||||
topk_weights=topk_weights,
|
w13_scale=layer.w13_weight_scale_bf16,
|
||||||
topk_ids=topk_ids,
|
w2=layer.w2_weight,
|
||||||
top_k=topk_ids.shape[1],
|
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)
|
return StandardCombineInput(hidden_states=output)
|
||||||
|
|
||||||
def apply_without_routing_weights(
|
def apply_without_routing_weights(
|
||||||
|
|||||||
Reference in New Issue
Block a user