已开启
feat: graph-based PTQ/QAT 支持 A8W4 量化 #269
QQR创建于 11 天前
feat: graph-based PTQ/QAT 支持 A8W4 量化 #269
已开启
合并受阻
atomgit-bot
11 天前 评论:
11 天前 评论:
变更摘要
本 PR 为 graph-based PTQ/QAT 增加 A8W4(INT4 权重 + INT8 激活)量化支持,目标算子为 conv2d/linear。核心改动包括:将 arq_retrain/ulq_scale_retrain 的 ONNX 导出收敛到新增的 add_qdq 助手,以原生 ONNX TensorProto.INT4(要求 opset 21)导出 4bit 权重;INT4 可 pack 性校验由原来的 Cin 轴改为最终 Deploy pack 轴(末轴,Gemm 且 transB=1 时为 axis 0);原先无法 pack 的层由「降级 INT8(A8W8)」改为「跳过该层、保持浮点」;Conv2dQAT/LinearQAT 开放 INT4 支持(LinearQAT 同时放开 channel_wise 限制),并在 qat_base.py 增加激活仅支持 per-tensor、INT4 权重须配 INT8 激活等约束;PackInt4WeightPass 改为直接写出原生 ONNX INT4,删除基于 numpy 的 nibble 手动打包逻辑。
主要改动
- 新增统一 QDQ 导出助手
add_qdq(新文件qdq_symbolic.py):将arq_retrain.py/ulq_scale_retrain.py中直接生成QUANTIZE_LINEAR/DEQUANTIZE_LINEAR的分支收敛为调用add_qdq;4bit 路径经_check_int4_export校验 opset 21 与原生TensorProto.INT4,per-channel 时校验wts_scales数量等于out_channels,Linear在 channel_wise 下用两次Transpose(perm=[1,0]) 将量化轴落于 axis 1。 - INT4 pack 轴校验从 Cin 轴改为最终 Deploy pack 轴:
check.py中check_int4_cin_pack_supported重构为check_int4_weight_pack_supported(原方法保留为兼容别名),不再因 grouped conv 或 Cin 为奇数拒绝,而是检查权重张量末轴(Gemm且transB=1时为 axis 0)长度是否为偶数,并在奇数时输出告警日志。 - 无法 pack 的 INT4 层由降级 INT8 改为跳过、保持浮点:
config_base.py的check_int4_weight_quant_axis与retrain_config_base.py的skip_int4_unpackable_layers(原downgrade_int4_unpackable_layers,保留兼容包装)不再将num_bits/dst_type从 4 改回 8,而是置quant_enable/retrain_enable为 False 并将该层从supported_layers与PARAM_POOL量化层列表剔除。 - QAT 模块开放 INT4 并增加配置约束:
Conv2dQAT/LinearQAT新增_supported_weight_dst_types = (INT8, INT4),LinearQAT.check_quantifiable移除原先「不支持 channel_wise」的RuntimeError;qat_base.py新增激活channel_wise报错、INT4 权重须搭配 INT8 激活的校验,且Conv2dQAT/LinearQAT分别在权重末轴(W)或out_features为奇数时抛出ValueError。 PackInt4WeightPass改为写出原生 ONNX INT4:删除基于 numpy 的pack_along_axis两两 nibble 打包逻辑,pack_int4_weight_node仅接收权重节点,校验末轴为偶数后通过TensorProtoHelper.set_data以INT4类型、保持原始 shape 写回权重(RNN 的 recurrence_weight 同样处理)。


atomgit-bot
11 天前 评论:
11 天前 评论:
11 天前 添加了label:cann-cla/yes
此处折叠了178条消息 查看更多
amct_pytorch/classic/graph_based/amct_pytorch/configuration/check.py
@@ -358,2 +356,3 @@
358356 dims = QuantOpInfo.get_node_tensor(wnode).dims
359- if cin_axis < len(dims) and dims[cin_axis] % 2 == 1:
357+ if dims and dims[pack_axis] % 2 == 1:
358+ actual_axis = pack_axis if pack_axis >= 0 else len(dims) + pack_axis
【一般】这里不需要处理负数,日志可以直接打印负数


amct_pytorch/classic/graph_based/amct_pytorch/configuration/check.py
@@ -363,0 +357,4 @@
357+ if dims and dims[pack_axis] % 2 == 1:
358+ actual_axis = pack_axis if pack_axis >= 0 else len(dims) + pack_axis
359+ LOGGER.logw(
360+ "Skip A8W4 layer '{}': ONNX weight shape {} has odd Deploy "
【一般】这里不一定是A8W4,不需要描述activation信息


amct_pytorch/classic/graph_based/amct_pytorch/configuration/check.py
@@ -337,2 +328,2 @@
337- """
338- # layer_name 预期能取到 node,取不到属于异常,交由 get_node_by_name 抛出
328+ def is_int4_weight_pack_axis_even(graph, layer_name):
329+ """Return whether all weights have an even final Deploy pack axis."""
【建议】补充函数功能描述,尤其是轴判断的信息


amct_pytorch/classic/graph_based/amct_pytorch/custom_op/arq_retrain/arq_retrain.py
@@ -52,2 +57,3 @@
52- Function: ArqRetrain foward funtion.
57+ Function: ArqRetrain forward function.
5358 """
59+ if is_dynamo_export():
AMCT内部没有效用dynamo=True的选项,也没有暴露接口给用户,这里的处理是不会被触发的?


amct_pytorch/classic/graph_based/amct_pytorch/custom_op/qdq_symbolic.py
@@ -0,0 +45,4 @@
45+def check_int4_export(wts_param):
46+ if not hasattr(TensorProto, 'INT4'):
47+ raise RuntimeError('INT4 export requires ONNX native TensorProto.INT4.')
48+ if _globals.GLOBALS.export_onnx_opset_version != 21:
torch.onnx.export 是否使用 opset 21 是AMCT自己控制的吧?
这里只是针对QAT单算子场景吗?


描述
支持 graph-based PTQ/QAT A8W4 量化,主要改动如下:
output_dtype=INT4且省略 zero-point;Conv2d per-channel 使用 axis=1,Linear 通过互逆 Transpose 使用 axis=1。transB=1检查轴 0,其他目标算子检查最后一轴。raw_data。Diff 摘要:
关联的Issue
https://gitcode.com/annqr/amct_open/issues/2
如何测试
git diff --check:PASS196 passed, 4 subtests passed执行命令:
文档更新
无文档更新。
类型标签