已合并
feat: 新增 MXFP4 QAT 伪量化 Linear(带 STE) #264
huangwuwei创建于 13 天前
feat: 新增 MXFP4 QAT 伪量化 Linear(带 STE) #264
已合并
Pull Request已成功合入, 合并人@CANN-robot
(感谢 huangwuwei 的贡献)13 天前 创建了 pull request,commit b391e6a3
atomgit-bot
13 天前 评论:
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.Function、MXFP4FakeQuantizer、MXFP4QATConfig)、linear.py(MXFP4QATLinear 与递归替换函数 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_mask在clip_grad=True时屏蔽饱和元素的反向梯度(clipped STE)。 - 算子后端自动切换:
mxfp4_quant_dequant的backend="auto"在 NPU 张量上按MXFP4_ASCENDC_PATH(缺省为同级../mxfp4_ascendc/python)自动加载 Ascend-C 内核,否则退回纯 PyTorch 路径;backend="npu"内核不可用或未编译时抛出带修复指引的RuntimeError。 - 新增可替换层与递归转换(
linear.py):MXFP4QATLinear继承nn.Linear,from_linear直接接管原Parameter对象而非拷贝(转换不额外占显存),convert_to_mxfp4_qat原地递归替换模型中所有nn.Linear,并支持skip_names按模块路径子串排除子树。 - 新增统一配置与无状态量化器:
MXFP4QATConfigdataclass 集中管理quantize_weight、quantize_input(W4A16/W4A4)、block_size、scale_factor、clip_grad、backend等参数;MXFP4FakeQuantizer为无参数、无 buffer 的 module 封装,使state_dict与 float 层完全一致,checkpoint 可与 float 模型双向加载。 - 文档更新:新增
mxfp4_qat/README.md(原理、API、接入自有训练框架的三种方式),并更新fakequant/README.md目录树与mxfp4_ascendc/README.md,说明推理侧(mxfp4_ascendc/)与训练侧(mxfp4_qat/)的分工。


atomgit-bot
13 天前 评论:
13 天前 评论:
13 天前 添加了label:cann-cla/yes
CANN-robot
13 天前 评论:
13 天前 评论:
此处折叠了43条消息 查看更多
yaoguangxiu
6 小时前 评论:
6 小时前 评论:
/approve


6 小时前 添加了label:approved
6 小时前 添加了label:lgtm
6 小时前 合入了pull request
描述
在
amct_pytorch/experimental/fakequant/下新增mxfp4_qat/,提供一个最小可用的 MXFP4 量化感知训练(QAT)实现:MXFP4QATLinear可直接替换torch.nn.Linear,前向按 MXFP4 数值伪量化,反向通过 STE(straight-through estimator)把梯度回传到高精度主权重,使模型在训练阶段就适应 MXFP4 的量化误差。改动原因:现有
mxfp4_ascendc/只解决推理侧的伪量化精度验证,缺少训练侧入口。对直转(PTQ)掉点较大的模型,需要 QAT 让权重"适应"MXFP4 后再导出,本 PR 补齐这一环。所采取的方法:
torch.autograd.Function。业界常见写法是x + (x_q - x).detach(),数值上等价;这里改成显式forward/backward,一是让梯度路径可读,二是为了能在backward里实现 clipped STE。clip_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,且MXFP4QATLinear是nn.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.pyautograd.Function、MXFP4FakeQuantizer、MXFP4QATConfigmxfp4_qat/linear.pyMXFP4QATLinear+convert_to_mxfp4_qat递归替换mxfp4_qat/README.md如何测试
前提条件:
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一行替换 / MegatronColumnParallelLinear等自定义 Linear 挂量化器 / 直接复用可微算子)、训练建议、限制说明。amct_pytorch/experimental/fakequant/README.md:目录树补充mxfp4_qat/,并说明mxfp4_ascendc/(推理侧精度验证)与mxfp4_qat/(训练侧)的分工。amct_pytorch/experimental/fakequant/mxfp4_ascendc/README.md:补充指向mxfp4_qat/的交叉引用。类型标签