AMCT FakeQuant 模拟伪量化工具包:MXFP4 推理精度验证与 QAT 训练实践
【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct
本文以 AMCT(昇腾 AI 处理器亲和的模型压缩工具)中的experimental/fakequant模块为主体,讲解当 NPU 尚不原生支持某类低比特格式时,如何在软件侧用“伪量化”复现 MXFP4(Microscaling FP4)的量化数值行为:既包括基于 Ascend-C 自定义 kernel 的推理侧精度验证算子,也包括基于 STE(直通估计器)的 MXFP4 量化感知训练(QAT)替换层。读完本文,读者可以独立完成 Ascend-C 算子的编译与正确性验证,并掌握在自有训练框架中接入 MXFP4 QAT Linear 的完整流程与调参要点。
1. 工具包定位:为什么需要模拟伪量化
在低比特格式落地过程中,常出现这样的时间差:算法侧已经需要评估 MXFP4 这类新格式的精度上限,但目标格式的 NPU 算子栈尚未合入,或当前环境无法启用对应的低比特计算单元。AMCT 的 FakeQuant 工具包(README)正是为此设计的:在硬件尚不原生支持某些量化格式时,在软件侧模拟对应格式的量化精度,以便进行精度验证与算法评估。
需要明确它的适用边界(原文档即有说明):
- 本模块属于试验特性(
experimental),接口与实现可能随硬件能力演进而调整; - 伪量化结果用于精度对齐与方案验证,不等同于目标硬件上真实低比特算子的性能表现;
- MXFP4 算子实现参考 amct_ops/hifloat8_cast 的三层结构(
op_kernel/op_extension/python),但因处于试验阶段暂不迁入amct_ops/。
工具包由两个职责互补的子模块组成:
fakequant/ ├── mxfp4_ascendc/ # MXFP4 Ascend-C 伪量化算子(目录对齐 amct_ops/hifloat8_cast) │ ├── op_kernel/ # device kernel + tiling │ ├── op_extension/# Torch host + TORCH_LIBRARY 注册 │ ├── python/mxfp4/# Python 包装 │ ├── reference/ # 纯 PyTorch 参考实现 │ └── tests/ ├── mxfp4_qat/ # MXFP4 量化感知训练(QAT):带 STE 的 nn.Linear 替代实现 │ ├── fake_quant.py# MXFP4 QDQ + STE autograd.Function + 量化器模块 │ └── linear.py # MXFP4QATLinear + convert_to_mxfp4_qat ├── README.md └── README_en.md其中mxfp4_ascendc/面向推理侧精度验证(快速跑出伪量化数值),mxfp4_qat/面向训练侧(让模型在训练中适应 MXFP4 误差);后者在 NPU 上会自动复用前者的算子加速。二者共享同一套 MXFP4 数值定义,保证“推理时看到的误差”和“训练时适应的误差”完全一致。
2. MXFP4 数值模型:E2M1 码本 + E8M0 逐块缩放
理解整个工具包的前提是 MXFP4 的数值模型:逐元素 FP4 E2M1 尾数 + 逐 block 的 E8M0(2 的幂)共享缩放。具体规则为:
- 最后一维每
block_size(默认 32)个相邻元素共享一个 scale; scale = 2^round(log2(max_abs / scale_factor)),scale_factor默认6.0,即把 block 内最大值映射到码本最大码值 6 附近;- 元素码本为
{0, ±0.5, ±1, ±1.5, ±2, ±3, ±4, ±6}(乘以 scale)。
这三条规则在仓库中有三份一一对应的实现,可以交叉印证:
(1)纯 PyTorch 参考实现reference/mxfp4_ref.py(L42-L61):先用F.pad对最后一维补齐到 block 整数倍(注意是按最后一维 padding 而非 flatten 后 padding,否则相邻行会错误地共享同一个 MXFP4 block),再按abs().amax求块内最大值,经clamp(min=2**-30)防下溢后用exp2(round(log2(...)))取 2 的幂作为 scale;元素舍入则通过一组torch.where(y_abs >= 中点, 码值, 原值)级联完成——中点序列(0.25, 0.75, 1.25, 1.75, 2.50, 3.50, 5.00)恰好是 E2M1 相邻码值的中点。
(2)QAT 侧的向量化实现mxfp4_qat/fake_quant.py(L115-L149):_block_scale()与参考实现逐步对应,元素舍入改用torch.bucketize(y_abs, midpoints, right=True)一步索引到最近码值(L132-L139),并缓存了按(device, dtype)组织的码本张量;文件头部常量BLOCK_SIZE = 32、SCALE_FACTOR = 6.0、MXFP4_E2M1_MAX = 6.0(L60-L62)定义了默认数值配置,最小块尺度_MIN_SCALE_RAW = 2**-30明确注释为“与mxfp4_tiling.h中的MXFP4_MIN_SCALE_RAW保持同步”。
(3)Ascend-C kernel 的宿主/设备共享常量op_kernel/mxfp4_tiling.h(L22-L27):
constexpr int32_t MXFP4_BLOCK_SIZE = 32; constexpr float MXFP4_SCALE_FACTOR = 6.0f; constexpr float MXFP4_INV_SCALE_FACTOR = 1.0f / 6.0f; constexpr float MXFP4_MIN_SCALE_RAW = 9.313225746e-10f; // 2^-30 constexpr int32_t MXFP4_BLOCKS_PER_TILE = 208; constexpr int32_t MXFP4_TILE_ELEMS = MXFP4_BLOCK_SIZE * MXFP4_BLOCKS_PER_TILE;kernel 内部以1/6.0为内置倒数因子,运行时再乘以一个可配置的invScaleMulBits(tiling 结构中保存的是 IEEE-754 位模式),这一设计正是 Python 包装层inv_scale_factor_scale参数的底层来源(见第 4 节)。
3. mxfp4_ascendc:NPU 侧 Ascend-C 伪量化算子
3.1 运行环境
子目录 README 给出的环境基线:
| 组件 | 版本 |
|---|---|
| 硬件 | Ascend 910B3 / 兼容 SoC |
| CANN | 8.2.RC1+ |
| Python | 3.10 (aarch64) |
| PyTorch | 2.6.0 |
| torch_npu | 2.6.0.post4 |
开源仓不附带预编译.so,需要按本机 SoC / CANN / Python ABI 自行编译。构建入口是 mxfp4_ascendc/CMakeLists.txt,其中定义了三个关键 CMake 缓存变量(L25-L28):
SOC_VERSION:默认Ascend910B3,可用SOC_VERSION=Ascend910_9392 bash build.sh形式覆盖;RUN_MODE:npu / cpu / sim;ASCEND_CANN_PACKAGE_PATH:默认/usr/local/Ascend/ascend-toolkit/latest。
从源码结构看,CMake 构建分两段:先经 CANN 的ascendc_kernel_cmake/ascendc.cmake用 ccec 编译 device kernel(op_kernel/mxfp4_kernel.cpp)为ascendc_kernels_${RUN_MODE}共享库,再编译 Torch 扩展libmxfp4_ops.so(L125-L149),后者链接 Torch、kernel 库与ascendcl,并显式探测torch._C._GLIBCXX_USE_CXX11_ABI以保持 ABI 一致——这就是“必须按本机 ABI 自行编译”的原因。
3.2 三层结构与 PyTorch 算子注册
目录布局对齐amct_ops/hifloat8_cast:
mxfp4_ascendc/ ├── op_kernel/ │ ├── mxfp4_kernel.cpp # Ascend-C device kernel │ └── mxfp4_tiling.h # Host/device 共用 tiling 常量与结构体 ├── op_extension/ │ ├── mxfp4_torch.cpp # PyTorch host:tiling + ACLRT_LAUNCH_KERNEL │ ├── ops.h # C++ host 接口声明(namespace AscendKernel) │ └── register.cpp # TORCH_LIBRARY_FRAGMENT(amct, ...) + Meta ├── python/ │ └── mxfp4/ │ ├── __init__.py # 加载 .so、自检、re-export │ └── ops.py # 薄 Python 包装(pad / dtype) ├── reference/ │ └── mxfp4_ref.py # 纯 PyTorch 参考实现 ├── CMakeLists.txt # 构建入口 ├── tests/ # 正确性 / inv_scale / benchmark └── README.md算子注册在 op_extension/register.cpp(L23-L45)中完成,值得注意的细节是注册方式:
TORCH_LIBRARY_FRAGMENT(amct, m) { m.def("quant_dequant_mxfp4(Tensor x, float inv_scale_factor_scale=1.0) -> Tensor"); }- 使用
TORCH_LIBRARY_FRAGMENT而非TORCH_LIBRARY,允许与 AMCT 其他扩展在同一命名空间amct下增量注册; - 通过
TORCH_LIBRARY_IMPL(amct, PrivateUse1, m)提供 NPU 设备实现(委托给AscendKernel::Mxfp4QuantDequantTorch,即 host 侧完成 tiling 并ACLRT_LAUNCH_KERNEL启动); - 通过
TORCH_LIBRARY_IMPL(amct, Meta, m)提供仅返回 shape 的 Meta 实现(at::empty(x.sizes(), x.options())),使该算子可安全穿过符号追踪/dynamo 类路径的 shape 推导。
因此底层原始调用形式为torch.ops.amct.quant_dequant_mxfp4(x_flat, 1.0),约定输入是float32、连续(contiguous)且 numel 为 32 的倍数的 flat 张量;AIV 核数由 host 侧PlatformAscendC::GetCoreNumAiv()运行时查询,无需手动指定。
3.3 Python 包装层的 padding 约定
python/mxfp4/ops.py(L25-L87)是面向用户的正式入口mxfp4.quant_dequant_mxfp4(x, block_size=32, inv_scale_factor_scale=1.0),它做了四件“脏活”:
- 校验
block_size == 32(NPU kernel 硬约束)与输入必须位于 NPU 设备; - 转为 float32 后先在最后一维 padding到 block 整数倍,再 reshape 成 flat——源码注释明确指出:若先 flatten 再 pad,当最后一维不是 32 的倍数时相邻行会合并进同一个 MXFP4 block,与参考实现不一致;
- 调用
torch.ops.amct.quant_dequant_mxfp4(x_flat, float(inv_scale_factor_scale)); - 恢复原 shape、切掉 padding、还原输入 dtype。
其中inv_scale_factor_scale与 QAT 侧的scale_factor参数是倒数关系:kernel 内置1/6.0,包装层用它乘上运行时乘子,故scale_factor = SCALE_FACTOR / inv_scale_factor_scale。这一换算在 fake_quant.py(L230-L236)中完成,用户无论走哪条路径,面对的语义参数都是scale_factor。
3.4 编译、正确性验证与性能
按子目录 README 的指引编译并测试(编译成功后.so会被 stage 到python/mxfp4/):
cd /path/to/mxfp4_ascendc # 编译(约数分钟);成功后自动将 .so 拷到 python/mxfp4/ bash build.sh # 测试正确性 + 性能 python tests/test_mxfp4.py # python tests/test_inv_scale.py # inv_scale 参数正确性 # python tests/bench_qdq.py # 额外性能对比正确性测试 tests/test_mxfp4.py(L35-L72)将 Ascend-C 输出与mxfp4_ref纯 PyTorch 参考实现逐元素比对(atol=1e-6),测试形状刻意覆盖了最后一维不是 32 倍数的场景——(2, 33)、(3, 17)、(4, 48, 33)——专门验证“按行 block”而非“flatten 后 pad”的语义;文档口径为与 PyTorch 参考实现 bit-exact 一致。
性能对比(同一文档口径,torch_npu 软件路径 vs Ascend-C 自定义 kernel):
| Shape | torch_npu | Ascend-C | 加速比 |
|---|---|---|---|
| (64, 4096) | 0.69 ms | 0.038 ms | 18.1x |
| (256, 4096) | 0.72 ms | 0.059 ms | 12.3x |
| (1024, 4096) | 0.72 ms | 0.219 ms | 3.28x |
即小矩阵约 18x、大矩阵约 3.3x 加速。该收益的解释与 QAT 侧的训练建议直接相关:torch_npu 软件路径每次 QDQ 要发射十几个 elementwise kernel,kernel 启动开销在小 shape 下占比极高,这正是自定义单 kernel 方案收益最大的区间。
3.5 快速使用
import sys sys.path.insert(0, "/path/to/mxfp4_ascendc/python") from mxfp4 import quant_dequant_mxfp4 x_npu = x.npu() result = quant_dequant_mxfp4(x_npu) # 等价底层调用(输入需已是 float32 flat,numel 为 32 的倍数) # result = torch.ops.amct.quant_dequant_mxfp4(x_flat, 1.0)API 签名:
quant_dequant_mxfp4( x: torch.Tensor, # 任意 shape,建议 float32,在 NPU 上 block_size: int = 32, # 量化 block 宽度(必须为 32) inv_scale_factor_scale: float = 1.0, ) -> torch.Tensor # 同 shape / dtype / device4. mxfp4_qat:MXFP4 量化感知训练
4.1 STE 与 clipped STE
量化是分段常量函数,几乎处处导数为 0,无法直接反传。STE(straight-through estimator,直通估计器)的做法是:前向用量化值、反向把量化算子当作恒等映射,从而让梯度穿过量化点抵达高精度权重:
forward : y = Q(x) backward: dL/dx = dL/dy # 普通 STE dL/dx = dL/dy * (|x| <= 6*scale) # clipped STE(clip_grad=True)clip_grad=True时把**发生截断(saturation)**的元素梯度置零。其动机来自 MXFP4 的数值特性:scale 被强制取整到 2 的幂(round(log2(...))可能向下取整),block 内最大值约有半数概率超出6*scale而被截断,这些位置的梯度方向具有误导性,屏蔽后训练通常更稳。
源码中该机制对应 fake_quant.py 的三处实现:
mxfp4_saturation_mask()(L239-L254):返回|x| > 6 * block_scale的 bool 掩码;_MXFP4FakeQuantSTE(L257-L272):显式torch.autograd.Function,forward中按需保存截断掩码,backward中grad_output.masked_fill(saturated, 0.0)后原样透传;- 普通路径
mxfp4_fake_quant()(L275-L296):可微的Tensor -> Tensor函数。
子目录 README 还说明了实现取向:它等价于 MindSpeed-LLM 的x + (x_q - x).detach()写法,但改用显式autograd.Function以便在backward里做梯度屏蔽。该模块只依赖torch(fake_quant.py的 import 仅torch与torch.nn.functional),可以单文件拷进任意训练仓使用。
4.2 MXFP4QATConfig 全参数
from mxfp4_qat import MXFP4QATConfig, MXFP4QATLinear, convert_to_mxfp4_qat| 字段 | 默认 | 说明 |
|---|---|---|
quantize_weight | True | 是否伪量化权重;置False则退化为纯激活实验 |
quantize_input | False | 是否伪量化层输入。False→ W4A16(推荐起点),True→ W4A4 |
block_size | 32 | 共享 scale 的元素数;Ascend-C 算子只支持 32 |
scale_factor | 6.0 | 增大 → scale 变小,inlier 分辨率更高但截断更多;减小则相反 |
clip_grad | False | True使用 clipped STE |
backend | "auto" | "auto"(NPU 上自动用 Ascend-C 算子,否则纯 PyTorch)/"torch"/"npu" |
dataclass 的__post_init__会对block_size、scale_factor的正数约束与backend取值做校验(fake_quant.py L335-L373)。
4.3 MXFP4QATLinear:无状态的 nn.Linear 替换
linear.py 中的MXFP4QATLinear是nn.Linear的子类(L33-L100),设计上有两个关键性质:
(1)主权重保持高精度,前向量化的是“一次性副本”。forward为F.linear(Q(x), Q(W), b)(bias 不量化):每次前向把当前高精度权重过一遍量化器得到伪量化值参与计算,optimizer 更新的始终是高精度“master weight”,loss 反映的却是 MXFP4 数值。
(2)state_dict与 float 层完全一致。由于量化器无参数、无 buffer,float 权重可以直接 load 进 QAT 模型;QAT 训练完的权重也可以 load 回 float 模型,或交给 AMCT 的 deploy 流程导出真实低比特权重。
构造方式:
layer = MXFP4QATLinear(in_features, out_features, config=MXFP4QATConfig()) layer = MXFP4QATLinear.from_linear(existing_linear, config) # 复用原 Parameter,不额外占显存from_linear(L71-L92)在meta设备上创建壳层后直接采用(adopt)原nn.Linear的weight/biasParameter 对象而非拷贝,因此不占额外显存,已引用这些 Parameter 的 optimizer 状态与 parameter group 也不受影响。
4.4 convert_to_mxfp4_qat:原地批量替换
convert_to_mxfp4_qat(module, config=None, skip_names=())(linear.py L103-L145)原地递归替换模块树下所有nn.Linear(含子类,因此 Megatron 风格的 Linear 子类也会被转换而非遗留 float)。要点:
skip_names按模块点分路径做子串匹配,命中则跳过该子树(连同其下所有层);- 已转换的
MXFP4QATLinear会被识别并跳过,因此重复调用是幂等 no-op; - 返回的是同一个 module 对象,仅用于链式调用。
convert_to_mxfp4_qat( model, MXFP4QATConfig(quantize_input=True, clip_grad=True), skip_names=("lm_head", "embed_tokens"), # 敏感层保持 float )4.5 快速开始
import sys sys.path.insert(0, ".../amct_pytorch/experimental/fakequant") from mxfp4_qat import MXFP4QATConfig, convert_to_mxfp4_qat model.load_state_dict(torch.load(ckpt)) # 从 float 预训练权重出发 convert_to_mxfp4_qat(model, MXFP4QATConfig(quantize_input=True)) # 其余训练代码不变4.6 接入自有训练框架的三种方式
方式一:模型里是标准nn.Linear。建模完成、加载完预训练权重之后,optimizer创建之前插入一行即可,其余训练代码不用改:
model = build_model() model.load_state_dict(torch.load(ckpt)) # 从 float 预训练权重出发 convert_to_mxfp4_qat(model, MXFP4QATConfig(quantize_input=True)) optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5) # ... 正常训练循环 ...from_linear复用原Parameter对象,所以在 optimizer 之后转换也不会失效;但放在 optimizer 之前更保险(避免 parameter group 引用悬空)。
方式二:框架有自定义 Linear(MegatronColumnParallelLinear等)。无法用继承替换时,把MXFP4FakeQuantizer(nn.Module形态,同样无状态)挂到层上、在forward里手动调用——这正是 MindSpeed-LLM 的做法:
from mxfp4_qat import MXFP4FakeQuantizer class FakeQuantColumnParallelLinear(ColumnParallelLinear): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.weight_quantizer = MXFP4FakeQuantizer() self.input_quantizer = MXFP4FakeQuantizer() def forward(self, input_, weight=None, **kwargs): input_ = self.input_quantizer(input_) # 父类 forward 读取 self.weight,因此临时替换 .data 后再恢复。 # 前向已在量化值上完成,梯度经 STE 正确回传到原始高精度权重。 original = self.weight.data self.weight.data = self.weight_quantizer(original) try: return super().forward(input_, weight=weight, **kwargs) finally: self.weight.data = originalMoE 的 GroupedMatmul 专家权重同理:在调用 GMM 之前对w1/w2和 permute 后的专家输入各过一次量化器。
方式三:只想复用量化算子。直接调用mxfp4_fake_quant(x),它是可微的Tensor -> Tensor函数,可放在任何位置(KV cache、logits、residual 等)。底层函数全集:
mxfp4_quant_dequant(x, block_size=32, scale_factor=6.0, backend="auto") # 无梯度,与 mxfp4_ascendc 参考实现 bit-exact mxfp4_fake_quant(x, block_size=32, scale_factor=6.0, clip_grad=False, backend="auto") # 带 STE mxfp4_saturation_mask(x, block_size=32, scale_factor=6.0) # 截断位置掩码 MXFP4FakeQuantizer(block_size=32, scale_factor=6.0, clip_grad=False, backend="auto") # nn.Module 形态4.7 训练建议(继承原文档)
- 从 float 预训练权重出发做 QAT 微调,不要从随机初始化开始训;
- 先 W4A16 再 W4A4:激活量化掉点通常明显大于权重量化,先确认
quantize_input=False能收敛; - 学习率取预训练的 1/10 左右并配 cosine 衰减;
- 敏感层保持 float:
lm_head、embedding、第一/最后一层通常通过skip_names排除; - NPU 上务必用 Ascend-C 后端:纯 PyTorch 路径每次 QDQ 有十几个 elementwise kernel,大模型训练开销不可忽略;
- 训练完成后用
amct_pytorch的 deploy 流程导出真实低比特权重;QAT 只是让权重“适应”MXFP4,导出仍需常规量化链路。
4.8 backend 解析机制:QAT 与 Ascend-C 算子的衔接
backend="auto"时的后端选择在 fake_quant.py(L224-L236)中实现:仅当张量位于 NPU 设备且Ascend-C kernel 可加载时才走 NPU 路径。kernel 包的查找顺序为:环境变量MXFP4_ASCENDC_PATH→ 同级../mxfp4_ascendc/python目录(_VENDORED_KERNEL_PATH,L81),成功与失败都会缓存,避免重复探测。算子需先自行编译(cd ../mxfp4_ascendc && bash build.sh);未编译或不在 NPU 上时自动退回纯 PyTorch 路径——两条路径结果 bit-exact 一致,仅速度不同;显式指定backend="npu"而算子不可用时,会抛出带修复指引(重新编译、设置MXFP4_ASCENDC_PATH或改用backend="torch")的RuntimeError。
5. 限制与使用前提
以下限制均继承自子目录 README,使用本工具包前应逐条确认:
- 属于试验特性(
experimental),接口可能调整; - QAT 侧只覆盖
nn.Linear;卷积、Embedding、Attention 内部的 matmul 未处理; - scale 与截断阈值均由数据静态推导,未实现可学习的 scale / clipping(LSQ、PACT 等);
block_size != 32只有纯 PyTorch 路径支持(Ascend-C kernel 与 Python 包装层均硬约束 32);- 伪量化仅复现 MXFP4 的数值行为,不代表目标硬件上真实低比特算子的性能;
- Ascend-C 算子需要 CANN 8.2.RC1+、aarch64 环境与对应 SoC,且开源仓不附带预编译
.so,需自行编译。
6. 小结
AMCT 的experimental/fakequant用一套统一的 MXFP4 数值定义串起了压缩工具链上的两个关键环节:mxfp4_ascendc提供与 PyTorch 参考实现 bit-exact 一致、并在小 shape 下最高约 18x 加速的 NPU 侧 QDQ 算子,支撑推理精度快速验证;mxfp4_qat在其上叠加 STE/clipped-STE 可微伪量化与MXFP4QATLinear,以“state_dict 与 float 完全兼容”为设计底线,让 QAT 可以低成本插入任意训练框架。由于二者共享同一数值模型(E2M1 码本 + E8M0 逐块 2 的幂 scale,常量在 mxfp4_tiling.h、mxfp4_ref.py、fake_quant.py 三处保持一致),训练阶段适应的误差与部署验证时观测的误差天然对齐,这正是“模拟伪量化”作为低比特格式落地前验证手段的价值所在。
【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考