Prhub

#49348 [ROCm][Quark][6/N] Use MXFP4 linear kernel abstraction for `aiter` backend

原始 PR 作者 fxmarty-amd 合并时间 2026-07-28 10:35 文件变更 6 提交数 19 评论 5 代码增减 +370 / -166

执行摘要

将 Quark OCP MX AITER 内核迁移至线性抽象层

标准化 mxfp4 线性后端,避免 quark_ocp_mx.py 特有的代码/调度,使所有检查点生产者(compressed-tensors 等)受益。具体见 PR body:'Motivation is to standardize mxfp4 linear backends throughout vLLM codebase, and avoid quark_ocp_mx.py-specific code/dispatch that can benefit all checkpoint producers'。

值得精读。该 PR 展示了在 vLLM 中注册一个新量化内核的完整流程:定义子类、实现抽象方法、注册到 _POSSIBLE_*_KERNELS_registered_linear_kernels,并在使用方通过 init_*_linear_kernel 获取实例。关注 __init__.py 中的注册模式和 test_mxfp4_kernel_selection.py 中 mock 平台枚举的测试技巧。

讨论亮点
  1. 命名编号调整:BowenBao 建议将本 PR 从 5/N 改为 6/N 并将 #48949 调整为 7/N,以反映 #46765 的合并顺序变更。已采纳。
  2. CI 测试失败处理:在 AMD CI 中观察到 test_gfx950_moe.py 失败,fxmarty-amd 指出该问题已在 #46491 中修复,不属于本 PR 的回归。
  3. GSM8K 阈值解释:fxmarty-amd 在 review 中说明 Qwen3-1.7B-MXFP4.yaml 中的 accuracy_threshold: 0.27 虽低,但与 Qwen3-0.6B-FP8.yaml 的阈值一致,且本 PR 的 GSM8K 结果与 main 分支完全匹配(0.2745),证明重构未引入精度退化。

实现拆解

  1. 新建文件 vllm/model_executor/kernels/linear/mxfp4/aiter.py:将原先在 quark_ocp_mx.py 内联注册的 gemm_with_dynamic_quantgemm_with_dynamic_quant_fake 自定义操作移植至此,并封装为 AiterMxfp4LinearKernel(继承 MxFp4LinearKernel),实现 is_supportedcan_implementprocess_weights_after_loadingapply_weights 四个抽象方法。
  2. 修改 __init__.py:导入 AiterMxfp4LinearKernel 并将其加入 _registered_linear_kernels['aiter'] 集合和 _POSSIBLE_MXFP4_KERNELS[PlatformEnum.ROCM] 列表。
  3. 重写 quark_ocp_mx.py:移除所有与 aiter 相关的导入、gemm_with_dynamic_quant 定义与注册代码;改为从 vllm.model_executor.kernels.linear 导入 init_mxfp4_linear_kernel,在 EmulateFalse 时调用 self.ocp_mx_linear = init_mxfp4_linear_kernel() 获取内核实例。
  4. 新增测试文件 tests/kernels/quantization/test_mxfp4_kernel_selection.py:验证抽象方法存在性、AITER 内核依赖 supports_mx 条件、初始化函数能正确分发到已注册内核、无匹配内核时抛 ValueError
  5. tests/quantization/test_quark.py 中新增 test_mxfp4_dynamic_quant_match_quark:对比 AITER 动态量化与 Quark 仿真 QDQ 的结果一致性。
  6. 新增 GSM8K 评测配置 Qwen3-1.7B-MXFP4.yaml,用于端到端回归验证重构前后精度不变。
文件 模块 状态 重要度
vllm/model_executor/kernels/linear/mxfp4/aiter.py 线性内核 added 9.36
vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py 量化方案 modified 8.48
tests/kernels/quantization/test_mxfp4_kernel_selection.py 内核选择 added 7.83
vllm/model_executor/kernels/linear/__init__.py 线性内核 modified 6.05
tests/quantization/test_quark.py 量化测试 modified 5.54
tests/evals/gsm8k/configs/Qwen3-1.7B-MXFP4.yaml 评测配置 added 4.3

关键符号

gemm_with_dynamic_quant gemm_with_dynamic_quant_fake AiterMxfp4LinearKernel.__init__ AiterMxfp4LinearKernel.is_supported AiterMxfp4LinearKernel.can_implement AiterMxfp4LinearKernel.process_weights_after_loading AiterMxfp4LinearKernel.apply_weights QuarkOCP_MX.__init__

关键源码片段

vllm/model_executor/layers/quantization/quark/schemes/quark_ocp_mx.py refactor

被重构的源文件:删除内联 AITER 注册代码,改为通过 init_mxfp4_linear_kernel 获取内核实例。

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM projectfrom vllm.logger import init_logger
from vllm.model_executor.kernels.linear import init_mxfp4_linear_kernelclass QuarkOCP_MX(QuarkScheme):
    def __init__(self, ...):
        # ...
        self.emulate = not current_platform.supports_mx() or (
            self.input_dtype != "mxfp4" or self.weight_dtype != "mxfp4"
        )
​
        # TODO: Move emulation code path as a kernel, and always
        # use init_mxfp4_linear_kernel.
        if not self.emulate:
            # 通过内核工厂获取 AiterMxfp4LinearKernel 实例
            self.ocp_mx_linear = init_mxfp4_linear_kernel()
        # ...

说明:修改后的 QuarkOCP_MX.__init__ 不再直接注册自定义 op,而是调用工厂函数 init_mxfp4_linear_kernel(),由框架根据平台和可用性选择最合适的 MxFp4LinearKernel 实现。

评论区精华

PR 编号调整为 6/N 设计

BowenBao 建议将本 PR 从 5/N 改为 6/N,并将 #48949 调整为 7/N,以反映 #46765 的合并顺序。

结论:已采纳,PR 标题和描述已更新。 · 已解决

CI test_gfx950_moe.py 失败 测试

AMD CI 中 test_gfx950_moe.py 失败,fxmarty-amd 指出失败已在 #46491 中修复,不是本 PR 引入。

结论:与 #46491 重复,无需本 PR 处理。 · 已解决

GSM8K 阈值设置说明 question

fxmarty-amd 在 review 中解释 accuracy_threshold 设为 0.27 虽然低,但与 Qwen3-0.6B-FP8 的阈值一致,且 GSM8K 测试通过。

结论:接受该解释。 · 已解决

风险与影响

风险较低,因为这是纯重构且数值结果已验证。但在 ROCm 平台上,如果 AITER 库不可用,init_mxfp4_linear_kernel 将抛出 ValueError(已在测试中覆盖),需确保部署环境中 AITER 已正确安装。此外,AiterMxfp4LinearKernelis_supported 要求 supports_mx() 为真,若平台误报可能导致运行时错误。建议通过 GSM8K 回归测试和单元测试持续验证。

对用户无感知,数值行为完全一致。对系统:ROCm 平台上的 MXFP4 推理路径从硬编码自定义操作迁移到统一内核选择框架,后续添加仿真内核(#48949)或其他后端变得标准化。对团队:降低了 quark_ocp_mx.py 的维护成本,与其他量化方案(如 NVFP4、FP8)代码结构对齐。

AITER 依赖路径变更 平台条件检查风险

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论