Repository navigation
[Refactor][Ops] Refactor MoE mlp and quantization methods - #13773
Conversation
Summary of ChangesHello, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request refactors the fused MoE MLP runtime to improve maintainability and extensibility. By abstracting the activation and quantization logic into separate components, the codebase removes redundant branching and allows for cleaner support of diverse quantization methods. This change is internal and does not impact user-facing APIs. Highlights
New Features🧠 You can now enable Memory (public preview) to help Gemini Code Assist learn from your team's feedback. This makes future code reviews more consistent and personalized to your project's style. Click here to enable Memory in your admin console. Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize the Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counterproductive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for GitHub and other Google products, sign up here. Footnotes
|
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
|
👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:
If CI fails, you can run linting and testing checks locally according Contributing and Testing. |
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[Ops][Feature] Refactor MoE activation and quantization methods for AscendSuggested PR Summary:
### What this PR does / why we need it?
This pull request refactors the MoE (Mixture of Experts) activation handling by introducing a modular `MoeMlpActivation` class and `get_moe_activation_method` helper. It also adds dynamic quantization helper methods (`_quant_hidden_states`, `apply_gmm1`, `apply_gmm1_act_quant`, and `apply_gmm2`) to the Ascend quantization schemes.
However, several critical issues were identified in the implementation:
- In `moe_comm_method.py`, `MoEFusedExpertsInput` (the class) is incorrectly used instead of the instance `fused_experts_input` to access the activation attribute.
- In `w8a8_dynamic.py`, helper functions `_require_single_tensor_for_swiglu_quant` and `cumsum_group_list` are used but not imported.
- Multiple bugs exist in `apply_gmm1`, `apply_gmm1_act_quant`, and `apply_gmm2` within `w8a8_dynamic.py`, including undefined variables (`hidden_states`), typos (`hidden_statesm`), missing return statements, and incorrect attribute access (`weight` instead of `weights`).
- The return type annotation for `apply_mlp` in `moe_activation.py` is incorrect as it returns a tuple instead of a single `Tensor`.
### Does this PR introduce _any_ user-facing change?
No.
### How was this patch tested?
No testing details were provided. It is highly recommended to fix the runtime NameErrors and attribute errors before testing.| token_dispatch_output=token_dispatch_output, | ||
| use_fusion_ops=self.use_fusion_ops, | ||
| ) | ||
| act_method = get_moe_activation_method(MoEFusedExpertsInput.activation) |
There was a problem hiding this comment.
MoEFusedExpertsInput is the class name, not the instance. Accessing MoEFusedExpertsInput.activation will fail at runtime. It should be changed to fused_experts_input.activation.
| act_method = get_moe_activation_method(MoEFusedExpertsInput.activation) | |
| act_method = get_moe_activation_method(fused_experts_input.activation) |
| from vllm_ascend.device.device_op import DeviceOperator | ||
| from vllm_ascend.distributed.parallel_state import get_mc2_group | ||
| from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input | ||
| from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEMlpComputeInput, build_fused_experts_input |
There was a problem hiding this comment.
The helper functions _require_single_tensor_for_swiglu_quant and cumsum_group_list are used in apply_gmm1_act_quant but are not imported, which will cause a NameError at runtime. They should be imported from vllm_ascend.ops.fused_moe.moe_mlp.
| from vllm_ascend.device.device_op import DeviceOperator | |
| from vllm_ascend.distributed.parallel_state import get_mc2_group | |
| from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input | |
| from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEMlpComputeInput, build_fused_experts_input | |
| from vllm_ascend.device.device_op import DeviceOperator | |
| from vllm_ascend.distributed.parallel_state import get_mc2_group | |
| from vllm_ascend.ops.fused_moe.moe_mlp import ( | |
| _require_single_tensor_for_swiglu_quant, | |
| cumsum_group_list, | |
| ) | |
| from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEMlpComputeInput, build_fused_experts_input |
| def apply_gmm1(self, mlp_compute_input: MoEMlpComputeInput): | ||
| hidden_states = torch_npu.npu_grouped_matmul( | ||
| x=[hidden_states], | ||
| weight=mlp_compute_input.weight.w1 | ||
| if isinstance(mlp_compute_input.weight.w1, list) | ||
| else [mlp_compute_input.weight.w1], | ||
| scale=[mlp_compute_input.weight.w1_scale[0].to(mlp_compute_input.weight.w2_scale[0].dtype)] | ||
| if isinstance(mlp_compute_input.weight.w1_scale, list) | ||
| else [mlp_compute_input.weight.w1_scale], | ||
| bias=None, | ||
| per_token_scale=[mlp_compute_input.weight.pertoken_scale], | ||
| split_item=2, | ||
| group_type=0, | ||
| group_list=mlp_compute_input.group_list, | ||
| output_dtype=torch.bfloat16, | ||
| )[0] |
There was a problem hiding this comment.
The apply_gmm1 method has several critical issues:
hidden_statesis referenced inx=[hidden_states]before it is defined. It should bemlp_compute_input.hidden_states.mlp_compute_input.weightis used instead ofmlp_compute_input.weights.mlp_compute_input.weight.pertoken_scaleis used, butMoEMlpComputeInputdoes not have aweightattribute or apertoken_scaleattribute under weights. It should usemlp_compute_input.dynamic_scale.- The method does not return
hidden_states.
| def apply_gmm1(self, mlp_compute_input: MoEMlpComputeInput): | |
| hidden_states = torch_npu.npu_grouped_matmul( | |
| x=[hidden_states], | |
| weight=mlp_compute_input.weight.w1 | |
| if isinstance(mlp_compute_input.weight.w1, list) | |
| else [mlp_compute_input.weight.w1], | |
| scale=[mlp_compute_input.weight.w1_scale[0].to(mlp_compute_input.weight.w2_scale[0].dtype)] | |
| if isinstance(mlp_compute_input.weight.w1_scale, list) | |
| else [mlp_compute_input.weight.w1_scale], | |
| bias=None, | |
| per_token_scale=[mlp_compute_input.weight.pertoken_scale], | |
| split_item=2, | |
| group_type=0, | |
| group_list=mlp_compute_input.group_list, | |
| output_dtype=torch.bfloat16, | |
| )[0] | |
| def apply_gmm1(self, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor: | |
| hidden_states = torch_npu.npu_grouped_matmul( | |
| x=[mlp_compute_input.hidden_states], | |
| weight=mlp_compute_input.weights.w1 | |
| if isinstance(mlp_compute_input.weights.w1, list) | |
| else [mlp_compute_input.weights.w1], | |
| scale=[mlp_compute_input.weights.w1_scale[0].to(mlp_compute_input.weights.w2_scale[0].dtype)] | |
| if isinstance(mlp_compute_input.weights.w1_scale, list) | |
| else [mlp_compute_input.weights.w1_scale], | |
| bias=None, | |
| per_token_scale=[mlp_compute_input.dynamic_scale], | |
| split_item=2, | |
| group_type=0, | |
| group_list=mlp_compute_input.group_list, | |
| output_dtype=torch.bfloat16, | |
| )[0] | |
| return hidden_states |
| def apply_gmm1_act_quant(self, mlp_compute_input: MoEMlpComputeInput): | ||
| hidden_states, pertoken_scale = self._quant_hidden_states( | ||
| mlp_compute_input.hidden_statesm, mlp_compute_input.dynamic_scale | ||
| ) | ||
| hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( | ||
| x=hidden_states, | ||
| weight=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weight.w1, name="w1"), | ||
| group_list=cumsum_group_list(mlp_compute_input.group_list, mlp_compute_input.group_list_type, 0), | ||
| weight_scale=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weight.w1_scale, name="w1_scale"), | ||
| x_scale=pertoken_scale, | ||
| bias=None, | ||
| use_mxfp_quant=False, | ||
| act_quant_type=torch.int8, | ||
| swiglu_limit=mlp_compute_input.swiglu_limit, | ||
| ) | ||
|
|
||
| return hidden_states, swiglu_out_scale |
There was a problem hiding this comment.
This method has two critical bugs:
mlp_compute_input.hidden_statesmhas a typo and should bemlp_compute_input.hidden_states.mlp_compute_input.weightis used instead ofmlp_compute_input.weights.
| def apply_gmm1_act_quant(self, mlp_compute_input: MoEMlpComputeInput): | |
| hidden_states, pertoken_scale = self._quant_hidden_states( | |
| mlp_compute_input.hidden_statesm, mlp_compute_input.dynamic_scale | |
| ) | |
| hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( | |
| x=hidden_states, | |
| weight=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weight.w1, name="w1"), | |
| group_list=cumsum_group_list(mlp_compute_input.group_list, mlp_compute_input.group_list_type, 0), | |
| weight_scale=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weight.w1_scale, name="w1_scale"), | |
| x_scale=pertoken_scale, | |
| bias=None, | |
| use_mxfp_quant=False, | |
| act_quant_type=torch.int8, | |
| swiglu_limit=mlp_compute_input.swiglu_limit, | |
| ) | |
| return hidden_states, swiglu_out_scale | |
| def apply_gmm1_act_quant(self, mlp_compute_input: MoEMlpComputeInput): | |
| hidden_states, pertoken_scale = self._quant_hidden_states( | |
| mlp_compute_input.hidden_states, mlp_compute_input.dynamic_scale | |
| ) | |
| hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant( | |
| x=hidden_states, | |
| weight=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weights.w1, name="w1"), | |
| group_list=cumsum_group_list(mlp_compute_input.group_list, mlp_compute_input.group_list_type, 0), | |
| weight_scale=_require_single_tensor_for_swiglu_quant(mlp_compute_input.weights.w1_scale, name="w1_scale"), | |
| x_scale=pertoken_scale, | |
| bias=None, | |
| use_mxfp_quant=False, | |
| act_quant_type=torch.int8, | |
| swiglu_limit=mlp_compute_input.swiglu_limit, | |
| ) | |
| return hidden_states, swiglu_out_scale |
| def apply_gmm2(self, mlp_compute_input: MoEMlpComputeInput, hidden_states, act_out_scale): | ||
| hidden_states = DeviceOperator.npu_grouped_matmul_gmm2( | ||
| hidden_states=hidden_states, | ||
| weight=mlp_compute_input.weight.w2, | ||
| weight_scale=mlp_compute_input.weight.w2_scale, | ||
| per_token_scale=act_out_scale, | ||
| group_list=mlp_compute_input.group_list, | ||
| group_list_type=mlp_compute_input.group_list_type, | ||
| input_dtype=None, | ||
| act_quant_type=torch.int8, | ||
| weight_quant_type=None, | ||
| scale_type=None, | ||
| per_token_scale_type=None, | ||
| use_bf16=True, | ||
| use_mxfp_quant=False, | ||
| bias=None, | ||
| fallback_output_dtype=mlp_compute_input.weight.w2_scale[0].dtype | ||
| if isinstance(mlp_compute_input.weight.w2_scale, list) | ||
| else mlp_compute_input.weight.w2_scale.dtype, | ||
| ) | ||
| return hidden_states |
There was a problem hiding this comment.
mlp_compute_input.weight is used instead of mlp_compute_input.weights.
def apply_gmm2(self, mlp_compute_input: MoEMlpComputeInput, hidden_states, act_out_scale):
hidden_states = DeviceOperator.npu_grouped_matmul_gmm2(
hidden_states=hidden_states,
weight=mlp_compute_input.weights.w2,
weight_scale=mlp_compute_input.weights.w2_scale,
per_token_scale=act_out_scale,
group_list=mlp_compute_input.group_list,
group_list_type=mlp_compute_input.group_list_type,
input_dtype=None,
act_quant_type=torch.int8,
weight_quant_type=None,
scale_type=None,
per_token_scale_type=None,
use_bf16=True,
use_mxfp_quant=False,
bias=None,
fallback_output_dtype=mlp_compute_input.weights.w2_scale[0].dtype
if isinstance(mlp_compute_input.weights.w2_scale, list)
else mlp_compute_input.weights.w2_scale.dtype,
)
return hidden_states|
|
||
| class MoeMlpActivation: | ||
| @staticmethod | ||
| def apply_mlp(mlp_compute_input: MoEMlpComputeInput, quant_method) -> torch.Tensor: |
There was a problem hiding this comment.
The return type annotation -> torch.Tensor is incorrect because the method returns a tuple of (hidden_states, before_gmm2_evt). It should be updated to tuple[torch.Tensor, torch.npu.Event].
| def apply_mlp(mlp_compute_input: MoEMlpComputeInput, quant_method) -> torch.Tensor: | |
| @staticmethod | |
| def apply_mlp(mlp_compute_input: MoEMlpComputeInput, quant_method) -> tuple[torch.Tensor, torch.npu.Event]: |
6f22bc8 to
f5ad9ea
Compare
f5ad9ea to
dba6efc
Compare
dba6efc to
1edc01c
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
| ) | ||
|
|
||
|
|
||
| def _gmm_swiglu_quant_fusion_enabled(use_mxfp_quant, fusion, dynamic_eplb, activation=None): |
There was a problem hiding this comment.
This function can be splitted. Only fusion and dynamic_eplb is useful. use_mxfp_quant and activation is useless after the refactoring.
| return int(num_experts - global_redundant_expert_num - num_shared_experts) | ||
|
|
||
|
|
||
| def _custom_gmm_swiglu_enabled(fusion, dynamic_eplb, activation=None): |
There was a problem hiding this comment.
This function can be splitted. Activation is useless after the refactoring.
|
|
||
| # ---- MLP gmm hooks (activation orchestrated by MoeActionMethod) ---- | ||
|
|
||
| def get_moe_weights(self, layer): |
There was a problem hiding this comment.
This function seems not consistent with the original code in AscendUnquantizedFusedMoEMethod's apply method
| ) | ||
| return final_hidden_states | ||
|
|
||
| def _get_mlp_weights(self, layer): |
There was a problem hiding this comment.
_get_mlp_weights processes both w1 and w2, but only one of them is used. Either in apply_gmm1 or apply_gmm2. Why not split it into two functions, one for w1, another for w2.
| w1_bias = layer.w13_bias if self.moe.has_bias else None | ||
| gate_up_out = torch_npu.npu_grouped_matmul( | ||
| x=[hidden_states], | ||
| weight=w1 if isinstance(w1, list) else [w1], |
There was a problem hiding this comment.
In the original unquant_apply_mlp, w1 is not a list. Maybe it's redundant to add the isinstance(w1, list)
| QuantType.W8A8MXFP, | ||
| QuantType.W4A4MXFP, | ||
| QuantType.W4A8MXFP, | ||
| QuantType.W4A16MXFP, |
There was a problem hiding this comment.
W4A16MXFP is not needed here.
|
|
||
| def apply_gmm1(self, mlp_compute_input: MoEMlpComputeInput): | ||
| # A16: activations stay unquantized; gmm1 dequantizes with antiquant. | ||
| hidden_states, _ = self._quant_hidden_states(mlp_compute_input.hidden_states, mlp_compute_input.dynamic_scale) |
There was a problem hiding this comment.
For specific quant types (e.g., here W4A16), it's better not to use the general _quant_hidden_states method. As you know for sure, hidden_states stay unquantized.
| supports_eplb = False | ||
| quant_type: QuantType = QuantType.W4A16MXFP | ||
| act_quant_type: torch.dtype | None = None | ||
| fused_activations = frozenset({"silu"}) |
There was a problem hiding this comment.
Consider an empty set for W4A16_MXFP4 (i.e., fused_activations = frozenset()). It has a similar logic as W4A16, the the following apply_gmm1_act_quant is not needed.
| output_dtype = w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype | ||
| return group_list, group_list_type, output_dtype | ||
|
|
||
| def get_moe_weights(self, layer): |
There was a problem hiding this comment.
It's better to use a more clear method name. get_moe_weights is a little confused. B.T.W, the final return branch seems for the case of non-fused mc2 (may be redundant)
| activation = mlp_compute_input.activation | ||
| act_name = getattr(activation, "value", activation) | ||
|
|
||
| if self.is_per_channel_weight and enable_custom_op() and act_name != "swigluoai_uninterleave": |
There was a problem hiding this comment.
It seems that act_name != "swigluoai_uninterleave" is not needed here.
| group_list_type=group_list_type, | ||
| swiglu_limit=mlp_compute_input.swiglu_limit, | ||
| ) | ||
| elif _custom_gmm_swiglu_enabled(mlp_compute_input.fusion, mlp_compute_input.dynamic_eplb, activation): |
There was a problem hiding this comment.
This branch is not needed for W4A8. In the original code, when use_w4a8_per_channel_gmm_swiglu is True, it only enter the first branch, which calls torch.ops._C_ascend.grouped_matmul_swiglu_quant_v2.
| ) | ||
| elif _gmm_swiglu_quant_fusion_enabled( | ||
| False, mlp_compute_input.fusion, mlp_compute_input.dynamic_eplb, activation | ||
| ): |
There was a problem hiding this comment.
Same as above, this branch is not needed for W4A8
| dispose_tensor(mlp_compute_input.hidden_states) | ||
| return hidden_states, swiglu_out_scale | ||
|
|
||
| def _soft_gmm1(self, mlp_compute_input: MoEMlpComputeInput, hidden_states: torch.Tensor, pertoken_scale): |
There was a problem hiding this comment.
The method name _soft_gmm1 seems confused. What does soft mean?
| scale=w2_scale, | ||
| bias=bias2, | ||
| per_token_scale=[act_out_scale], | ||
| split_item=2, |
There was a problem hiding this comment.
Can use torch_npu.npu_grouped_matmul directly.
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
f9303cd to
01697b3
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
01697b3 to
c759197
Compare
e3eb9a7 to
403532c
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
1 similar comment
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
403532c to
27ef236
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
27ef236 to
f11e5c4
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
…methods - Remove moe_mlp.py; distribute quant_apply_mlp/unquant_apply_mlp into per-quant-method gmm hooks (apply_gmm1/apply_gmm1_act_quant/apply_act_quant/apply_gmm2) under quantization/methods - Introduce MoeActionMethod per activation family (silu default, gelu, gelu_tanh, swiglustep, swigluoai_uninterleave) orchestrating the MLP stage with the quant method injected at runtime - Merge the redundant MC2/non-MC2 fused gmm+swiglu+quant branches; keep the int32+npu_dequant_swiglu_quant path for swigluoai_uninterleave - Extract A5 MXFP kernel branches from device_op.py into the respective FusedMoEMethod implementations - Thread layer through MoEFusedExpertsInput/MoEMlpComputeInput so apply() no longer passes weight tensors; FUSED_MC2 builds weights via quant_method.get_fused_mc2_weights(layer) - Keep the 310P legacy path untouched - Align MoE MLP UTs with the functional orchestration and the get_fused_mc2_weights rename: drop the orphan TestSwigluOaiDynamicMxQuant, mock DeviceOperator.clipped_swiglu for swigluoai_uninterleave, and track fused_w1/w2_scale_bias in the fused-MC2 weight test Signed-off-by: earthmanylf <yulinfeng2@huawei.com>
f11e5c4 to
ff41910
Compare
|
/rerun Rerun (failed jobs only):
|
What this PR does / why we need it?
Come from #13318 and #13220
Right now, moe_mlp has too many branches about different quantization method and different activations. In this PR, we split moe_mlp's methods into differnet quantization method.
Does this PR introduce any user-facing change?
No. This is an internal refactor of the fused MoE runtime contract and its test coverage.
How was this patch tested?
The following models are tested and can be curled successfully:
A3:
float: qwen30b-a3b
w8a8: qwen30b-a3b-w8a8
w4a8: kimi2.6-w4a8
megamoe, dynamic eplb: DeepSeek-V4-Flash-w8a8-mtp
A5:
w8a8mxfp: step3.7-w8a8mxfp
w4a8mxfp: qwen3-30b-w4a8mxfp
w4a4mxfp: glm5.1-w4a4mxfp