Prhub

#22352 :sparkles: [llm][npu][quant] Add W8A8 MXFP8 quantization support for Qwen3 Dense on Ascend NPU

原始 PR 作者 TallMessiWu 合并时间 2026-06-16 14:45 文件变更 8 提交数 41 评论 29 代码增减 +376 / -11

执行摘要

为 Ascend NPU Qwen3 Dense 模型添加 W8A8 MXFP8 量化支持

参考PR body:"closes part of the NPU quantization gap tracked in issue #21584"。目标是支持Ascend NPU A5系列硬件的MXFP8量化,以降低内存带宽和计算开销,同时保持接近无损精度。

值得精读,尤其process_weights_after_loading中对.contiguous()的差异化处理是NPU量化的关键设计决策,显示了平台硬件特性对软件实现的约束。

讨论亮点
  • 预转置与block-scale映射:gemini-code-assist建议离线量化权重预转置时不调用.contiguous(),作者确认已实现并解释原因。
  • ImportError容忍:ping1jing2问为什么跳过fused_rope_qk_mqa导入错误,作者说明缺失kernel会导致模块导入失败进而使ModelRegistry静默回退,已改为try/except+warning。
  • weight dtype告警:ping1jing2建议对非预期weight dtype添加warning,作者接受并添加。
  • 方案委托kernel:TamirBaydasov建议ModelSlimMXFP8Scheme应委托处理逻辑给NPUMXFP8LinearMethod,作者同意但计划在后续PR中重构。

实现拆解

  1. NPU旁路与分发:在fp8.pyFp8Config.get_min_capability()中增加NPU返回0的旁路,避免CUDA capability检查;在get_quant_method()中增加NPU+use_mxfp8分支,分发到NPUMXFP8LinearMethod
  2. 在线量化线性方法:在linear_method_npu.py新增NPUMXFP8LinearMethod类。create_weights以原始FP16/BF16 dtype预留权重;process_weights_after_loading通过npu_dynamic_mx_quant在线量化为MXFP8并预转置为[in, out]布局(不调用.contiguous()以保留分析视图);apply对激活逐token量化后调用npu_quant_matmul(group_sizes=[1,1,32])。
  3. 离线量化方案:在modelslim_mxfp8.py新增ModelSlimMXFP8Scheme类,处理msmodelslim预量化权重(float8_e4m3fn + uint8 scale)。create_weights创建对应dtype的参数;process_weights_after_loading转置但不调用.contiguous()以保持block-scale映射;apply_weights类似在线路径进行激活量化和矩阵乘法。
  4. 注册与导入调整:在modelslim/modelslim.pyget_linear_scheme表中注册("W8A8_MXFP8", ModelSlimMXFP8Scheme);在schemes/__init__.py中先导入基类再导入新scheme以避免循环依赖。
  5. 导入修复:在rotary_embedding/base.py中将fused_rope_qk_mqa导入包裹在try/except ImportError中,并将符号置为None,避免缺失kernel导致整个模块导入失败进而使模型静默回退。
  6. 文档:更新quantization.mdx添加mxfp8行;更新ascend_npu_quantization.mdx补充LLM dense使用示例。
文件 模块 状态 重要度
python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8.py 量化方案 added 9.28
python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py 量化实现 modified 8.69
python/sglang/srt/layers/rotary_embedding/base.py 位置编码 modified 6.61
python/sglang/srt/layers/quantization/fp8.py 量化框架 modified 5.73
python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py 量化方案 modified 5.8
python/sglang/srt/layers/quantization/modelslim/modelslim.py 量化注册 modified 4.89
docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_quantization.mdx 文档 modified 3.63
docs_new/docs/advanced_features/quantization.mdx 文档 modified 3.28

关键符号

NPUMXFP8LinearMethod.create_weights NPUMXFP8LinearMethod.process_weights_after_loading NPUMXFP8LinearMethod.apply ModelSlimMXFP8Scheme.create_weights ModelSlimMXFP8Scheme.process_weights_after_loading ModelSlimMXFP8Scheme.apply_weights Fp8Config.get_min_capability Fp8Config.get_quant_method ModelSlimConfig.get_linear_scheme

关键源码片段

python/sglang/srt/layers/rotary_embedding/base.py dependency-wiring

修复缺失 fused_rope_qk_mqa kernel 时模块导入失败的问题,通过 try/except 容忍缺失,并在 forward_npu 中检查 None 后回退。

if _is_npu:
    import torch_npu
​
    # `fused_rope_qk_mqa` is an optional fast-path kernel shipped with
    # `sgl_kernel_npu`. Older NPU CANN / sgl_kernel_npu builds may not
    # include it. If we let the ImportError propagate, importing this
    # module fails, which in turn causes `ModelRegistry` to silently skip
    # every model that depends on it (and fall back to HF Transformers
    # without quantisation awareness — see PR #22352). We tolerate the
    # missing kernel so model loading still works; call sites must check
    # for `None` and use the generic rope path. A warning is emitted so
    # the missing kernel is visible in logs instead of being silently
    # swallowed.
    try:
        from sgl_kernel_npu.norm.fused_rope_qk_mqa import fused_rope_qk_mqa
    except ImportError:
        fused_rope_qk_mqa = None
        logger.warning(
            "sgl_kernel_npu.norm.fused_rope_qk_mqa is unavailable; "
            "falling back to the generic rope implementation.  Upgrade "
            "sgl_kernel_npu to enable the fused kernel."
        )# ... inside forward_npu:
        if (
            fused_rope_qk_mqa is not None
            and query.shape[0] * query.shape[1] < 65535
        ):
            return fused_rope_qk_mqa(
                query,
                key,
                cos_sin,
                self.rotary_dim,
                self.is_neox_style,
            )
        else:
            return self.forward_native(positions, query, key, offsets)

评论区精华

预转置权重和 scale 时不调用 .contiguous() 以保持 block-scale 映射 正确性

gemini-code-assist 建议在离线量化路径中将权重和 scale 预转置并避免 .contiguous() 以免破坏 block-scale 映射。

结论:作者确认已按建议实现,在 modelslim_mxfp8.py 中使用 .data.transpose() 而不调用 .contiguous()。 · 已解决

缺失 fused_rope_qk_mqa 时跳过 ImportError 的原因 正确性

ping1jing2 问为什么跳过 ImportError,作者解释缺失 kernel 会导致模块无法导入进而使 ModelRegistry 静默回退到 HF Transformers 失去量化意识。

结论:作者使用 try/except 并将 fused_rope_qk_mqa 置为 None,同时发出 warning。 · 已解决

在线量化路径中应添加 weight dtype 告警 设计

ping1jing2 建议在 NPUMXFP8LinearMethod.process_weights_after_loading 中对非预期 weight dtype 添加 warning。

结论:作者接受并在 6f0e711 中添加了 logger.warning。 · 已解决

ModelSlimMXFP8Scheme 应委托 kernel 给 NPUMXFP8LinearMethod 设计

TamirBaydasov 指出 scheme 的职责应仅限于 weight 创建,处理逻辑应委托给 kernel。

结论:作者同意但计划在后续 PR 中重构,以免影响当前合并。 · acknowledged

风险与影响

  1. NPU条件导入:虽然使用current_platform.is_npu()保护,但若平台检测出错可能导致CUDA/CPU上错误导入torch_npu而崩溃。
  2. 布局差异风险:在线与离线路径对.contiguous()的处理不同(在线调用、离线不调用),若后续修改未注意此差异可能引入难以调试的数值错误。
  3. rotary_embedding修改影响面:该修改影响所有NPU模型加载,若try/except未覆盖其他缺失kernel场景可能导致静默回退降低推理精度。
  4. 缺少测试覆盖:没有直接添加单元测试或端到端测试,回归风险依赖NPU CI。

用户:Ascend NPU A5用户可通过--quantization mxfp8或msmodelslim预量化权重获得MXFP8加速;非NPU用户无影响。
系统:新引入NPUMXFP8LinearMethodModelSlimMXFP8Scheme,代码量适中,与原量化框架集成良好。
团队:后续需关注已识别的重构项(委托kernel、统一torch.ops.npu),并补充测试。

核心路径变更 (rotary_embedding) 缺少直接测试覆盖 NPU 专用代码条件保护需验证

关联 Issue

#21584 [RFC][NPU] Ascend NPU A5 Support for MXFP8/MXFP4 Quantization

完整报告

参与讨论