news 2026/9/18 19:21:00

CANN AMCT Conv2dQAT 量化感知训练算子 API 实战指南:从构造、配置到源码原理

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN AMCT Conv2dQAT 量化感知训练算子 API 实战指南:从构造、配置到源码原理

CANN AMCT Conv2dQAT 量化感知训练算子 API 实战指南:从构造、配置到源码原理

【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct

导读

Conv2dQAT 是 CANN AMCT(昇腾 AI 处理器亲和的模型压缩工具仓)提供的 2D 卷积量化感知训练(QAT)单算子,用于将浮点torch.nn.Conv2d替换为带量化感知训练能力的算子,在训练/重训过程中学习数据的截断上下限与量化因子,从而显著降低 INT8/INT4 量化带来的精度损失。本文以官方 API 文档为主体,结合仓库源码与测试用例,系统讲解 Conv2dQAT 的两种构造方式、全部参数语义、量化配置项(retrain_data_config / retrain_weight_config)、底层量化流程(IFMR 初始化 → ULQ/ARQ 重训)与常见约束,帮助你在昇腾场景下快速完成卷积层的 QAT 接入。

产品支持情况

Conv2dQAT 在以下昇腾硬件产品上获得支持(来源于 Conv2dQAT.md):

产品是否支持
Ascend 950PR / Ascend 950DT
Atlas A3 训练系列产品 / Atlas A3 推理系列产品
Atlas A2 训练系列产品 / Atlas A2 推理系列产品

功能说明

Conv2dQAT 用于构造 Conv2d 的 QAT 算子。与普通卷积层不同,该算子在网络前向中会依次完成:

  1. 激活量化:对输入 activation 先通过 IFMR(Initialization For MinMax Range)算法在初始化阶段统计量化范围,再经由 ULQ(截断上下限)重训算法学习/微调 clip_min、clip_max 等参数;
  2. 权重量化:对卷积权重执行 ARQ(Adaptive Rounding-based Quantization)或 ULQ 重训量化;
  3. 浮点卷积计算:使用量化后的激活与量化后的权重调用F.conv2d完成卷积,bias 保持浮点直接参与计算。

从源码结构看,Conv2dQAT定义于 conv2d.py,它同时继承torch.nn.Conv2d与统一的 QAT 基类QATBase(见 qat_base.py),因此在具备完整 Conv2d 行为的同时,自动获得量化参数注册、IFMR 初始化、ULQ/ARQ 重训等 QAT 能力。

函数原型

Conv2dQAT 提供两种等价的构造方式:

  • 直接构造接口
qat = amct_pytorch.nn.module.quantization.conv2d.Conv2dQAT(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode, device, dtype, config)
  • 基于原生算子构造接口(推荐用于已有浮点模型的改造):
qat = amct_pytorch.nn.module.quantization.conv2d.Conv2dQAT.from_float(mod, config)

其中amct_pytorch.nn.module.quantization.conv2d是公共转出模块,内部将Conv2dQAT从实现路径 re-export 出来(见 conv2d.py),两种写法最终指向同一个类。

参数说明

表 1:直接构造接口参数

参数名输入/输出说明
in_channels输入含义:输入 channel 个数。数据类型:int
out_channels输入含义:输出 channel 个数。数据类型:int
kernel_size输入含义:卷积核大小。数据类型:int/tuple
stride输入含义:卷积步长。数据类型:int/tuple;默认值:1
padding输入含义:填充大小。数据类型:int/tuple;默认值:0
dilation输入含义:kernel 元素之间的间距。数据类型:int/tuple;默认值:1
groups输入含义:输入和输出的连接关系。数据类型:int;默认值:1
bias输入含义:是否开启偏置项参与学习。数据类型:bool,其他数据类型(比如整数、字符串、列表等)按照 Python 真值判断规则转换;默认值:True
padding_mode输入含义:填充方式。使用约束:仅支持zeros
device输入含义:运行设备。默认值:None
dtype输入含义:torch 数值类型。torch 数据类型,仅支持 torch.float32
config输入含义:量化配置。数据类型:dict;默认值:None(不传时按默认配置执行 QAT,详见下文)

表 2:基于原生算子构造接口参数

参数名输入/输出说明
mod输入含义:待量化的原生 Conv2d 算子。数据类型:torch.nn.Module(必须是torch.nn.Conv2d,否则抛 TypeError)
config输入含义:量化配置。数据类型:dict;默认值:None

from_float的转换逻辑位于 qat_base.py:首先校验mod必须是_float_module(即torch.nn.Conv2d)的实例;随后从_required_params(in_channels、out_channels、kernel_size、stride、padding、dilation、groups、bias、padding_mode)中提取原生算子的超参数;bias会被转换为"是否存在偏置"的布尔值;最后构造 QAT 算子并直接复用原生算子的 weight 与 bias 参数,保证转换前后模型权重完全一致。

config 量化配置

config为 dict 类型,官方参考样例如下(详见 Conv2dQAT.md):

config = { "retrain_enable": True, "retrain_data_config": { "dst_type": "INT8", "batch_num": 10, "fixed_min": False, "clip_min": -1.0, "clip_max": 1.0 }, "retrain_weight_config": { "dst_type": "INT8", "weights_retrain_algo": "arq_retrain", "channel_wise": False } }

各配置项的完整语义请参见 量化配置参数说明,核心要点归纳如下。

retrain_enable
  • 作用:该层是否进行量化感知训练。
  • 类型:bool;取值范围:true 或 false。
  • 说明:true表示该层进行 QAT(默认行为);false表示该层不进行量化感知训练,此时前向中激活与权重直接透传(见 qat_base.py)。
  • 推荐配置:true;可选参数。
retrain_data_config(数据/激活量化配置)

类型为 dict,包含以下可选参数:

  • batch_num:量化使用的 batch 数量。类型 int;取值范围大于 0;默认值 1。batch_num * batch_size为量化使用的校准集图片数量(batch_size 为每个 batch 所用的图片数量),建议校准集图片数量不超过 50 张。
  • clip_max:截断量化算法上限。类型 float;要求clip_max > 0。若配置则固定算法截断上限;若不配置,则通过 IFMR 算法学习获取上限。推荐取 activation 分布最大值 max 的0.3*max ~ 1.7*max区间。
  • clip_min:截断量化算法下限。类型 float;要求clip_min < 0。若配置则固定算法截断下限;若不配置,则通过 IFMR 算法学习获取下限。推荐取 activation 分布最小值 min 的0.3*min ~ 1.7*min区间。
  • fixed_min:数据量化算法下限固定开关。类型 bool;true表示固定下限且下限为 0,false表示不固定下限。默认不选。
  • dst_type:量化位宽类型。类型 string;当前激活量化支持 INT8/INT16,默认 INT8。

配置示例(来自官方文档):

"retrain_data_config": { "dst_type": "INT8", "batch_num": 10, "fixed_min": False, "clip_min": -1.0, "clip_max": 1.0 }
retrain_weight_config(权重量化配置)

类型为 dict,包含以下可选参数:

  • weights_retrain_algo:权重量化算法。类型 string;取值范围ulq_quantize(ULQ 截断上下限量化算法)与arq_retrain(ARQ 量化算法),默认arq_retrain。从源码看,实际算法分发键为arq_retrain/ulq_retrain(见 qat_base.py)。
  • channel_wise:是否对每个 channel 采用不同的量化因子。类型 bool;true表示每个 channel 独立量化、量化因子不同;false表示所有 channel 共享量化因子。默认 true(推荐)。
  • dst_type:量化位宽类型。类型 string;当前仅支持 INT8,默认为 INT8。从源码结构看,Conv2dQAT 的_supported_weight_dst_types = (INT8, INT4)(见 conv2d.py),即权重同时支持 INT4 量化(INT4 权重要求激活为 INT8,且权重的 W 轴宽度为偶数)。

返回值说明

  • 直接构造:返回构造的 QAT 单算子实例(Conv2dQAT)。
  • 基于原生算子构造:返回torch.nn.Module转化后的 QAT 单算子(仍是Conv2dQAT实例,且沿用了原生算子的 weight/bias)。

调用示例

示例一:直接构造

from amct_pytorch.nn.module.quantization.conv2d import Conv2dQAT Conv2dQAT(in_channels=1, out_channels=1, kernel_size=1, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode='zeros', device=None, dtype=None, config=None)

示例二:基于原生算子构造(模型改造场景)

import torch from amct_pytorch.nn.module.quantization.conv2d import Conv2dQAT conv2d_op = torch.nn.Conv2d(in_channels=1, out_channels=1, kernel_size=1, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode='zeros', device=None, dtype=None) Conv2dQAT.from_float(mod=conv2d_op, config=None)

示例三:携带量化配置的完整使用(可运行)

参考 test_qat_op.py 中的用法,可以构造一个带配置的 Conv2dQAT 并执行前向:

import torch from amct_pytorch.classic.graph_based.amct_pytorch.nn.module.quantization.conv2d import Conv2dQAT quant_config = { "retrain_enable": True, "retrain_data_config": { "dst_type": "INT8", "batch_num": 3, "fixed_min": False, "clip_min": -1.0, "clip_max": 1.0, }, "retrain_weight_config": { "dst_type": "INT8", "weights_retrain_algo": "arq_retrain", "channel_wise": True, }, } qat_conv = Conv2dQAT(in_channels=3, out_channels=16, kernel_size=1, stride=1, padding=0, config=quant_config) inputs = torch.randn((3, 3, 224, 224)) # 4 维输入:N, C, H, W output = qat_conv.forward(inputs) print(output.shape)

底层实现原理

类定义与继承关系

Conv2dQAT(nn.Conv2d, QATBase)(见 conv2d.py)同时继承原生torch.nn.Conv2dQATBase

  • 继承nn.Conv2d获得卷积超参数与 weight/bias 参数管理能力;
  • 继承QATBase获得统一的量化感知训练实现(IFMR 初始化、ULQ/ARQ 重训、量化参数注册、Dynamo 导出支持等)。

__init__中先以原生方式初始化nn.Conv2d,随后调用QATBase.__init__(self, 'Conv2d', device=device, config=config),其中'Conv2d'为层类型标识,用于后续统计/分发。

前向计算流程

forward(见 conv2d.py)要求输入必须是 4 维(N, C, H, W),流程为:

  1. forward_qat(inputs)返回量化后的激活quantized_acts与量化后的权重quantized_wts
  2. 调用F.conv2d(quantized_acts, quantized_wts, bias, stride, padding, dilation, groups)完成卷积,bias 以浮点形式直接参与。

forward_qat(见 qat_base.py)的核心逻辑为:

  • 输入 dtype 必须是torch.float32,否则抛 ValueError;
  • retrain_enable=True时:若尚未完成初始化(do_init=True),调用acts_quant_init用 IFMR 模块统计首个 batch 的 scale/offset/clip 范围;否则调用acts_quant走 ULQ 重训;随后调用wts_quant走 ARQ/ULQ 权重量化;
  • retrain_enable=False时:激活与权重直接透传,等价于普通浮点卷积。

量化参数注册

_register_qat_params(见 qat_base.py)会为算子注册以下可训练/缓冲参数:

  • 激活侧:acts_clip_maxacts_clip_min(可训练 Parameter,默认 1.0 / -1.0)、acts_scaleacts_offset_deployacts_clip_max_preacts_clip_min_precur_batch
  • 权重侧:wts_scaleswts_offsets(数量由 channel_wise 决定,True时为 out_channels 个,False时为 1 个)、wts_offsets_deploys_rec_flag

这些参数会在重训过程中由copy_tensor持续更新,最终用于部署阶段的 Q/DQ 节点导出(Dynamo ONNX 导出时通过add_qdq_dynamo/add_weight_qdq_dynamo构建,见 qat_base.py)。

配置校验与使用约束

QATBase._check_qat_config(见 qat_base.py)会在构造时严格校验配置,常见约束包括:

  • 激活dst_type仅支持 INT8/INT16;
  • 激活量化仅支持 per-tensor(channel_wise必须为 False,否则报错 "Activation quantization only supports per-tensor");
  • batch_num必须是大于 0 的整数;
  • fixed_min必须是 bool;
  • clip_min必须是小于 0 的 float,clip_max必须是大于 0 的 float;
  • 权重量化算法仅支持arq_retrain/ulq_retrain
  • INT4 权重量化要求激活为 INT8;
  • padding_mode 仅支持 'zeros'check_quantifiable(见 conv2d.py)在开启重训且 padding_mode 非 zeros 时抛 ValueError;同时若权重为 INT4 且权重形状 W 轴宽度为奇数也会报错。

测试用例验证

仓库测试 test_qat_op.py 对 Conv2dQAT 覆盖了完整的行为验证,可作为接入时的自测参考:

  • test_conv2d_qat_from_float_success/test_conv2d_qat_from_float_failed_padding_mode_not_zeros:验证from_float成功路径,以及 padding_mode 为reflect时抛 ValueError;
  • test_conv2d_qat_from_float_failed_ori_op_not_conv2d:传入Conv3d抛 TypeError;
  • test_conv2d_qat_forward/test_conv2d_qat_forward_ulq_retrain:验证默认 ARQ 与 ULQ 重训配置下的前向均可正常输出(输入(3, 3, 224, 224));
  • test_conv2d_qat_unsupport_shape_inputs:非 4 维输入抛 RuntimeError;
  • test_conv2d_qat_accepts_int4_per_tensor_and_per_channel:验证 INT4 权重在 per-tensor(1 个 scale)与 per-channel(4 个 scale)下的参数注册数量;
  • test_conv2d_qat_int4_odd_kernel_width_raises:INT4 权重 W 轴宽度为奇数时报错;
  • test_grouped_conv2d_qat_int4_even_width_is_supported:验证 groups=4 的分组卷积同样受支持。

测试中还提供了多组quant_configs组合(如ulq_retrain+batch_num=3fixed_min=Trueclip_min/clip_max同时配置等),可用于覆盖不同算法分支的前向验证。

常见问题与最佳实践

  1. padding_mode 必须为 'zeros':构造或from_float时若传入reflectreplicatecircular等填充方式,在开启重训时直接抛 ValueError,请在替换前先调整原生模型。
  2. 输入必须为 4 维且 dtype 为 float32:Conv2dQAT 前向只接受 (N, C, H, W) 的 torch.float32 输入,否则分别抛 RuntimeError / ValueError。
  3. config 传 None 时的默认行为:从源码看,config 为 None 时按空 dict 处理,retrain_enable默认 True,即默认开启量化感知训练;clip_min/clip_max不配置时通过 IFMR 在首轮前向自动学习范围。
  4. clip_min/clip_max 建议成对配置:需要手动固定截断范围时,应同时给出clip_min < 0clip_max > 0的合理取值(官方推荐基于 activation 分布的 0.3~1.7 倍区间),并合理设置batch_num使校准数据量适中。
  5. 模型改造流程:遍历named_modules()找到torch.nn.Conv2d实例,逐一用Conv2dQAT.from_float(module, config=...)替换,随后按常规流程训练/重训,最后借助 AMCT 的部署能力导出带 Q/DQ 节点的量化模型。更多量化配置细节请参考 量化配置参数说明。

【免费下载链接】amctAMCT是CANN提供的昇腾AI处理器亲和的模型压缩工具仓。项目地址: https://gitcode.com/cann/amct

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 19:20:37

Spark Job aborted与stage failure排查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/18 19:17:15

YOLOv26不是新模型:RK3588部署前必须厘清的商用模型本质

1. Yolov26 是什么&#xff1f;先别急着部署&#xff0c;得搞清它到底是不是“真新模型” 看到标题里那个 Yolov26 &#xff0c;我第一反应是——等等&#xff0c;YOLO 系列目前公开的主流版本是 YOLOv8、YOLOv9、YOLOv10&#xff08;2024 年中已开源&#xff09;&#xff0…

作者头像 李华
网站建设 2026/9/18 19:14:45

从零实现LTC细胞:液态神经网络核心单元手写指南

1. 项目概述&#xff1a;为什么LTC细胞值得从零手写一遍&#xff1f;液态神经网络&#xff08;Liquid Time-Constant Networks, LTN&#xff09;这几年在时序建模领域悄悄火了起来&#xff0c;尤其在低功耗边缘设备、生物信号处理、实时控制系统这些对延迟敏感、资源受限的场景…

作者头像 李华
网站建设 2026/9/18 19:14:36

MySQL查询语句全解析:从SELECT *到索引优化与排错实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华