已合并
feat: 新增 MXFP4 QAT 伪量化 Linear(带 STE) #264
huangwuwei创建于 13 天前
feat: 新增 MXFP4 QAT 伪量化 Linear(带 STE) #264
已合并
huangwuwei创建于 13 天前
huangwuwei
13 天前

描述

amct_pytorch/experimental/fakequant/ 下新增 mxfp4_qat/,提供一个最小可用的 MXFP4 量化感知训练(QAT)实现MXFP4QATLinear 可直接替换 torch.nn.Linear,前向按 MXFP4 数值伪量化,反向通过 STE(straight-through estimator)把梯度回传到高精度主权重,使模型在训练阶段就适应 MXFP4 的量化误差。

改动原因:现有 mxfp4_ascendc/ 只解决推理侧的伪量化精度验证,缺少训练侧入口。对直转(PTQ)掉点较大的模型,需要 QAT 让权重"适应"MXFP4 后再导出,本 PR 补齐这一环。

所采取的方法

  • STE 用显式 torch.autograd.Function。业界常见写法是 x + (x_q - x).detach(),数值上等价;这里改成显式 forward/backward,一是让梯度路径可读,二是为了能在 backward 里实现 clipped STE。
  • 支持 clipped STEclip_grad=True)。因为 block scale 被强制取整到 2 的幂且可能向下取整,block 内最大值有相当概率超出 6*scale 而被截断,这些位置的梯度方向具有误导性,屏蔽后训练更稳。
  • 算子后端自动切换backend="auto" 在 NPU 张量上自动复用本仓已有的 Ascend-C 算子(按 MXFP4_ASCENDC_PATH → 同级 ../mxfp4_ascendc/python 查找),否则退回纯 PyTorch,两条路径数值 bit-exact 一致。算子未编译或显式指定 backend="npu" 不可用时,抛出带修复指引的 RuntimeError
  • state_dict 与 float 层完全一致。量化器无参数、无 buffer,且 MXFP4QATLinearnn.Linear 子类,因此 float checkpoint 可直接 load 进 QAT 模型,QAT 权重也能 load 回 float 模型交给 deploy 链路。from_linear 接管原 Parameter 对象而非拷贝,转换不额外占显存。
  • fake_quant.py 仅依赖 torch,可作为单文件拷进第三方训练框架(Megatron / DeepSpeed / HF Trainer)使用。
  • 码本查表用 torch.bucketize(阈值恰为码本中点,落在第几个区间即码本下标)而非多次链式 torch.where,减少 elementwise kernel 数量,对训练热路径有意义。

新增文件:

文件 作用
mxfp4_qat/fake_quant.py MXFP4 QDQ、STE autograd.FunctionMXFP4FakeQuantizerMXFP4QATConfig
mxfp4_qat/linear.py MXFP4QATLinear + convert_to_mxfp4_qat 递归替换
mxfp4_qat/README.md 原理、API、接入自有训练框架的三种方式

如何测试

前提条件torch(CPU 即可)。NPU 后端为可选项,需先在 mxfp4_ascendc/ 执行 bash build.sh 编译算子。

最小验证:

import sys, torch
sys.path.insert(0, "amct_pytorch/experimental/fakequant")
from mxfp4_qat import MXFP4QATConfig, MXFP4QATLinear, convert_to_mxfp4_qat

model = torch.nn.Sequential(torch.nn.Linear(64, 32))
convert_to_mxfp4_qat(model, MXFP4QATConfig(quantize_input=True))
assert isinstance(model[0], MXFP4QATLinear)

model(torch.randn(4, 64)).sum().backward()
assert torch.isfinite(model[0].weight.grad).all()

代码检查pre-commit run --files <本次改动文件> 全部 12 个 hook 通过(含 OAT 开源合规、ruff check/format、codespell)。

文档更新

  • 新增 amct_pytorch/experimental/fakequant/mxfp4_qat/README.md:MXFP4 与 STE 原理、目录结构、快速开始、完整 API 表、接入自有训练框架的三种方式(标准 nn.Linear 一行替换 / Megatron ColumnParallelLinear 等自定义 Linear 挂量化器 / 直接复用可微算子)、训练建议、限制说明。
  • 更新 amct_pytorch/experimental/fakequant/README.md:目录树补充 mxfp4_qat/,并说明 mxfp4_ascendc/(推理侧精度验证)与 mxfp4_qat/(训练侧)的分工。
  • 更新 amct_pytorch/experimental/fakequant/mxfp4_ascendc/README.md:补充指向 mxfp4_qat/ 的交叉引用。

类型标签

likedislike
Pull Request已成功合入, 合并人@CANN-robot
(感谢 huangwuwei 的贡献)
Hhuangwuwei
13 天前 创建了 pull request,commit b391e6a3
atomgit-bot
atomgit-bot
13 天前 评论:

变更摘要

本 PR 在 amct_pytorch/experimental/fakequant/ 下新增 mxfp4_qat/ 模块,提供一个最小可用的 MXFP4 量化感知训练(QAT)实现:MXFP4QATLinear 作为 torch.nn.Linear 的直接替换,前向按 MXFP4 数值伪量化、反向通过 STE(straight-through estimator)把梯度回传到高精度主权重,补齐了现有 mxfp4_ascendc/ 推理侧验证所缺的训练侧入口。核心内容包括 fake_quant.py(MXFP4 QDQ、STE torch.autograd.FunctionMXFP4FakeQuantizerMXFP4QATConfig)、linear.pyMXFP4QATLinear 与递归替换函数 convert_to_mxfp4_qat)以及配套文档更新。

主要改动

  • 新增 MXFP4 伪量化与 STE 实现fake_quant.py):mxfp4_quant_dequant 按 block(block_size=32)做 QDQ——每 block 共享 2 的幂 E8M0 scale,尾数用 torch.bucketize 按码本中点阈值查 E2M1 码表;_MXFP4FakeQuantSTE 用显式 torch.autograd.Function 实现前向伪量化、反向直通,并支持通过 mxfp4_saturation_maskclip_grad=True 时屏蔽饱和元素的反向梯度(clipped STE)。
  • 算子后端自动切换mxfp4_quant_dequantbackend="auto" 在 NPU 张量上按 MXFP4_ASCENDC_PATH(缺省为同级 ../mxfp4_ascendc/python)自动加载 Ascend-C 内核,否则退回纯 PyTorch 路径;backend="npu" 内核不可用或未编译时抛出带修复指引的 RuntimeError
  • 新增可替换层与递归转换linear.py):MXFP4QATLinear 继承 nn.Linearfrom_linear 直接接管原 Parameter 对象而非拷贝(转换不额外占显存),convert_to_mxfp4_qat 原地递归替换模型中所有 nn.Linear,并支持 skip_names 按模块路径子串排除子树。
  • 新增统一配置与无状态量化器MXFP4QATConfig dataclass 集中管理 quantize_weightquantize_input(W4A16/W4A4)、block_sizescale_factorclip_gradbackend 等参数;MXFP4FakeQuantizer 为无参数、无 buffer 的 module 封装,使 state_dict 与 float 层完全一致,checkpoint 可与 float 模型双向加载。
  • 文档更新:新增 mxfp4_qat/README.md(原理、API、接入自有训练框架的三种方式),并更新 fakequant/README.md 目录树与 mxfp4_ascendc/README.md,说明推理侧(mxfp4_ascendc/)与训练侧(mxfp4_qat/)的分工。
likedislike
atomgit-bot
atomgit-bot
13 天前 评论:

代码审查

✅ 未发现问题

likedislike
CANN-robotCANN-robot成员
13 天前 添加了label:cann-cla/yes
CANN-robot
CANN-robot成员
13 天前 评论:

CLA Signature Pass

huangwuwei, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
此处折叠了43条消息 查看更多
yaoguangxiu成员
6 小时前 评论:

/approve

likedislike
CANN-robotCANN-robot成员
6 小时前 添加了label:approved
fujun19成员
6 小时前 评论:

/lgtm

likedislike
CANN-robotCANN-robot成员
6 小时前 添加了label:lgtm
CANN-robotCANN-robot成员
6 小时前 合入了pull request