执行摘要
- 一句话:为 NPU A5 上的 Qwen3 MoE 模型添加 W8A8 MXFP8 量化
- 推荐动作:该 PR 设计周密,值得仔细阅读,特别是以下方面:融合路由量化如何避免额外的量化 kernel 启动、在线/离线方法如何共享同一套
NPUMXFP8MoEMethod 内核,以及 create_moe_runner 中通过 isinstance 检查 moe_runner_backend 以防止错误分发。同时,审查反馈中的迭代过程(继承 vs 组合、内核提取、DeepEP 支持)展示了良好的代码演化实践,可供团队借鉴。
功能与动机
在 Ascend NPU A5 上实现 MXFP8 量化可以显著提升推理性能和内存效率。此 PR 是 issue #21584 中规划的一部分,该 issue 概述了在 NPU A5 上支持各种模型系列(扩散、LLM、VLM)的 MXFP8/MXFP4 量化。具体而言,此 PR 实现了 Qwen3 MoE 模型的 W8A8 MXFP8 支持,是之前密集模型 W8A8(PR #22352)和 W4A4(PR #23795)的延续。
实现拆解
- 配置分发:在
Fp8Config.get_quant_method() 中为 NPU 添加 FusedMoE 分支,返回新的 NPUMXFP8OnlineMoEMethod,实现在线量化入口。
- 在线 MoE 方法:创建
NPUMXFP8OnlineMoEMethod(继承 UnquantizedFusedMoEMethod),仅重写 create_moe_runner,将 NPUMXFP8MoEMethod 实例作为 w13_kernel / w2_kernel 附加到层上,其余权重创建和后处理逻辑继承自未量化方法。
- 核心 MoE 方法:实现
NPUMXFP8MoEMethod(继承 _NPUMoEMethodBase),包含权重在线量化(_quantize_weight_online 使用 npu_dynamic_mx_quant 对 3D 权重进行量化)、权重后处理(转换为 FRACTAL_NZ 布局 via npu_format_cast)以及前向传播。前向中,gmm1 使用 GroupedMatmulSwigluQuant(融合 gate+up+swiglu+requant,返回量化后的激活和块缩放),gmm2 使用 GroupedMatmul。
- 离线 MoE Scheme:创建
ModelSlimMXFP8MoEScheme(继承 ModelSlimMoEScheme),加载预量化的 float8_e4m3fn 权重和 uint8 块缩放,并在 process_weights_after_loading 中委托给 NPUMXFP8MoEMethod 进行布局转换。
- 融合路由量化:在
npu_moe_init_routing_v2 中通过 quant_mode=3 在路由时一同完成激活量化,避免额外调用 npu_dynamic_mx_quant。通过 _normalize_mxfp_scale 将扁平缩放因子转换为 pair-split 布局以匹配后续内核需求。
- 路由门量化修复:离线检查点中
mlp.gate 的权重量化现在通过 quant_model_description.json 描述驱动,而不是硬编码为 BF16,从而正确应用 MXFP8 缩放因子。
- 动态量化封装:将
hidden_states_quant.py 重命名为 quant.py,并让 HiddenStatesDynamicQuant 根据 dtype 选择 npu_dynamic_mx_quant(MXFP8)或 npu_dynamic_quant(INT8/INT4),使模块可复用。
- 文档:更新了
ascend_npu_quantization.mdx 和 quantization.mdx,描述 MXFP8 MoE 的使用方法。
- 测试:PR 未包含直接对应的测试文件,仅通过端到端启动脚本验证。
关键文件:
python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8_moe.py(模块 量化方案;类别 source;类型 data-contract;符号 ModelSlimMXFP8MoEScheme, init, create_weights, process_weights_after_loading): 新增文件,实现离线 MXFP8 MoE scheme,定义权重创建和布局转换委托。
python/sglang/srt/hardware_backend/npu/quantization/moe_methods.py(模块 MoE方法;类别 source;类型 dependency-wiring;符号 _require_e8m0_dtype, NPUMXFP8MoEMethod, init, _quantize_weight_online): 核心文件,实现 NPUMXFP8MoEMethod,包含在线权重量化、前向传播以及量化和非量化权重的分支处理。
python/sglang/srt/hardware_backend/npu/moe/matmul.py(模块 内核封装;类别 source;类型 dependency-wiring;符号 GroupedMatmulSwigluQuant, forward): 新增 GroupedMatmulSwigluQuant 类,封装 fused gate+up+swiglu+requant 内核,供 MXFP8 及未来 block-scaled 方案复用。
关键符号:ModelSlimMXFP8MoEScheme.init, ModelSlimMXFP8MoEScheme.create_weights, ModelSlimMXFP8MoEScheme.process_weights_after_loading, NPUMXFP8MoEMethod.init, NPUMXFP8MoEMethod._quantize_weight_online, NPUMXFP8MoEMethod.process_weights_after_loading, NPUMXFP8MoEMethod.apply, NPUMXFP8MoEMethod.apply_fused_gmm1_swiglu, NPUMXFP8OnlineMoEMethod.init, NPUMXFP8OnlineMoEMethod.create_moe_runner, GroupedMatmulSwigluQuant.forward, _normalize_mxfp_scale, HiddenStatesDynamicQuant.init
关键源码片段
python/sglang/srt/hardware_backend/npu/moe/matmul.py
新增 GroupedMatmulSwigluQuant 类,封装 fused gate+up+swiglu+requant 内核,供 MXFP8 及未来 block-scaled 方案复用。
from abc import ABC, abstractmethod
from typing import Tuple
import torch
class BaseMatmul(ABC):
@abstractmethod
def forward(self, ...):
pass
class GroupedMatmul(BaseMatmul):
# ... 已有实现 ...
class GroupedMatmulSwigluQuant(BaseMatmul):
"""Grouped matmul with swiglu and requantisation fused into one kernel.
Used for the gate/up projection (gmm1) of block-scaled MoE: the kernel
returns (quantized_activations, block_scale) instead of a single tensor.
Unlike ``GroupedMatmul`` it takes no ``output_dtype`` — output dtype comes
from ``quant_dtype`` in ``scale_args``.
"""
def forward(
self,
layer: torch.nn.Module,
weight_prefix: str,
hidden_states: torch.Tensor,
expert_tokens: torch.Tensor,
output_dtype: torch.dtype = None,
group_list_type: int = 1,
transposed: bool = True,
**scale_args,
) -> Tuple[torch.Tensor, torch.Tensor]:
weight = getattr(layer, f"{weight_prefix}_weight", None)
if weight is None:
raise AttributeError(
f"Weight attribute '{weight_prefix}_weight' not found in layer"
)
# This op requires a CUMULATIVE group_list; the dispatcher produces
# COUNT form (group_list_type=1). Convert here.
group_list = (
expert_tokens.cumsum(0) if group_list_type == 1 else expert_tokens
)
return torch.ops.npu.npu_grouped_matmul_swiglu_quant_v2(
x=hidden_states,
weight=[weight] if transposed else [weight.transpose(1, 2)],
group_list=group_list,
**scale_args,
)
评论区精华
审查中主要讨论了以下设计决策:
风险与影响
关联脉络
- PR #21584 [RFC][NPU] Ascend NPU A5 Support for MXFP8/MXFP4 Quantization: 总体跟踪 issue,本 PR 是其 MoE 部分的具体实现。
- PR #22352 :sparkles: [llm][npu][quant] Add W8A8 MXFP8 quantization support for Qwen3 Dense on Ascend NPU: 密集模型 W8A8 支持,本 PR 复用了其基础设施(Fp8Config 分发、NPUMXFP8LinearMethod、ModelSlimMXFP8Scheme)。
- PR #23795 :sparkles: [llm][npu][quant] Add W4A4 MXFP4 quantization support for Qwen3 Dense on Ascend NPU: 密集模型 W4A4 支持,与本 PR 同属 NPU 量化系列,共享底层 utils。
- PR #25663 [Refactor] Unified Ascend MoE execution stack (AscendRunnerCore + dispatcher + moe_methods): 上游重构 PR,本 PR 在合并时进行了适配,将原有 fused_moe_method_npu.py 中的逻辑移植到新的 moe_methods / runner 结构中。
参与讨论