这一篇,我想跟各位认真聊聊量化感知训练(QAT)这件事。先交代一下背景:我之前在一个边缘计算设备上部署了一个分类模型,PyTorch 训练完浮点精度 92.3%,直接转 INT8 之后掉到 84.7%,整整掉了接近 8 个点,而且排查了半天不是后处理的问题,就是量化本身的锅。后来用 QAT 重新走了一遍,精度拉回 91.6%,基本无损。这篇文章就把整个思路、操作流程和踩过的坑整理出来,希望能帮到正在被“INT8 精度下降”折磨的朋友。
1. 为什么好好的模型一转 INT8 就掉点——量化误差的根源
很多人拿到“精度下降”的第一反应是:是不是我转换姿势不对?是不是校准集选得不好?其实大部分情况下,校准集原因只占一小部分,本质问题出在“量化”这个动作本身对模型权重和激活值分布造成的破坏上。
1.1 量化的本质:用 256 个整数格子描述海量浮点数值
INT8 量化说白了就是把原来 FP32 的浮点数值,映射到 [-128, 127](或者 [0, 255])这 256 个整数格子上。你可以把它想象成一张 4096 级灰阶的高清图,强行压成 256 级灰阶:整体轮廓还在,但细节层次一定会有损失。
这个映射关系里最关键的两个参数是 scale(缩放系数)和 zero_point(零点偏移),它们共同决定了一个浮点区间如何映射到整数区间。PyTorch 在转换时通过 observer(观察器)统计每一层权重和激活的数值范围(min/max 或百分位),然后算出一个最优的 scale 和 zero_point。问题就出在这个“数值范围”的统计上——它只是训练集的一部分数据,不是全量数据,更不保证推理时输入的分布和它一致。
1.2 三种最典型的掉点模式,看看你属于哪种
我归纳了一下,实际项目里掉点基本逃不出这三种情况:
第一种是“离群值撕裂分布”。权重或者激活里如果有个别特别大的数值(比如超过 99.99% 分位的异常激活),observer 为了覆盖这个最大值,会把 scale 拉大,导致大部分正常数值的量化步长也变大,精度掉得就特别厉害。这在 Transformer、BERT 这类模型里尤其常见——Softmax 之后或者 Attention 层里经常出现极端激活值。
第二种是“逐层误差累积放大”。量化误差不是独立存在的。第一层激活的量化误差会传递到第二层,第二层再叠加自己的误差……在深层网络里,误差会像滚雪球一样越来越大。尤其是 BN 层、残差连接这种对数值敏感的结构,误差经过相加之后往往会成倍放大。
第三种是“校准集和部署场景分布不一致”。你拿着 ImageNet 的验证集去校准,实际部署时输入的是监控摄像头拍到的画面——光照、角度、噪声全都不一样,量化参数自然就不准。这种情况最坑,因为模型本身没毛病,纯粹是量化参数“没见过世面”。
1.3 哪些模型最容易掉点,哪些相对没事
我把过去两年帮别人排查的项目做了个简单归类,不一定绝对,但大方向是准的:
| 模型类型 | PTQ 掉点情况 | 原因 |
|---|---|---|
| 大 CNN(ResNet50/101 等) | 通常 0.5%~3% | 层数深,但结构规整,误差相对可控 |
| 轻量 CNN(MobileNet、ShuffleNet) | 3%~8%,甚至更多 | 深度可分离卷积对量化极敏感,数值范围碎片化 |
| Transformer(BERT、ViT) | 2%~10% | Softmax、LayerNorm、残差结构容易放大误差 |
| 检测/分割模型(YOLO、DeepLab) | 2%~8% | 多任务头、特征金字塔导致逐层误差累加复杂 |
| RNN/LSTM | 很容易崩 | 循环结构内误差反复叠加,时序依赖敏感 |
如果你手头的模型 PTQ 只掉 0.5 个点,说实话没必要上 QAT,用一些校准技巧就能救回来;但如果你掉 3 个点以上,尤其是轻量网络或者 Transformer,老老实实走 QAT 是性价比最高的方案。
2. QAT 为什么能拯救精度:从“被动接受误差”到“主动适应误差”
PTQ 是模型已经训练完了,再拿一批数据去统计量化参数,模型本身对量化这件事一无所知。QAT 的思路反过来:让模型在训练的时候就“提前体验”量化带来的误差,迫使权重适应一个更粗糙的数值环境。用一个不太严谨但很好理解的类比——QAT 就像让一个习惯看高清水牌的人先戴一副磨砂眼镜去背招牌,等他摘掉眼镜的时候,反而能更准确地认出低分辨率下的字。
2.1 伪量化节点(FakeQuantize)到底做了什么
QAT 的核心是 FakeQuantize 模块。它的作用是在前向传播的时候,把浮点数值模拟成“量化后再反量化”的结果——先量化到整数(模拟精度损失),再反量化回浮点(保证后续计算还是浮点)。这样一来,模型在训练时的前向计算就带上了真实的量化噪声。
关键是反向传播。量化函数本身不可导(它是个阶梯函数,几乎处处导数为 0),所以 PyTorch 里用了直通估计器(STE,Straight-Through Estimator)来处理:梯度直接“穿透”量化函数,绕过去回传到上游。这样做虽然数学上不严格,但实践下来收敛效果很好。这也是 QAT 能在算力不爆炸的前提下完成训练的根本原因。
2.2 QAT、PTQ 和 Fine-tuning 的区别,一张表说清楚
很多人分不清 QAT 和普通 fine-tuning,甚至有人问“我量化后直接拿数据再训几轮不就行了”。这里我把三者的区别梳理一下:
| 对比维度 | PTQ(训练后量化) | QAT(量化感知训练) | 普通 Fine-tuning |
|---|---|---|---|
| 是否修改网络结构 | 否 | 是(插入 FakeQuantize 节点) | 否 |
| 是否需要数据 | 少量校准集即可 | 需要训练集(或足够有代表性的数据) | 需要训练集 |
| 训练过程是否模拟量化 | 不模拟 | 前向模拟量化误差 | 不模拟 |
| 精度恢复能力 | 弱(只能调 scale 和 zero_point) | 强(权重主动适应量化噪声) | 弱(模型不知道量化的存在) |
| 耗时 | 分钟级 | 小时级(GPU 训练) | 取决于数据量 |
这么看就很明显了:QAT = 量化噪声 + 微调训练 + 权重自适应,三者缺一不可。
2.3 为什么 Fuse(融合)在 QAT 里是必经之路
做 QAT 的时候,你会发现官方教程里第一步永远是做模块融合——最常见的是把 Conv + BN + ReLU 融合成一个 Conv。原因有二:第一,融合之后可以减少一个量化边界,避免中间激活值被量化一次然后又反量化一次,这种重复操作会带来额外误差;第二,BN 层在推理时本来就会被吸收进卷积权重里,QAT 阶段先融合,训练时模拟的数值行为和最终部署的行为更一致。
这块一定要注意:如果没有融合直接跑 QAT,得到的量化模型精度往往比融合后跑 QAT 要低不少,而且这个差距在轻量网络上会被放大。后面实操部分我会具体演示怎么融合。
3. PyTorch QAT 全流程实操:从模型改造到 INT8 导出
接下来是动手环节。我用 torchvision 自带的 ResNet18 和一个简单的水果分类数据集做演示,完整跑一遍 QAT 流程。环境是 PyTorch 2.1 + CUDA 11.8,如果你用的是 CPU 版本,流程完全一样,只是训练慢一些。
3.1 环境准备与模型改造
首先装依赖,PyTorch 本身自带量化工具,不需要额外安装其他库:
pip install torch torchvision然后改造模型。QAT 要求模型知道自己“哪些层是输入、哪些层是输出”,以便插入量化节点。最直接的方式是在模型的 forward 开头加 QuantStub,结尾加 DeQuantStub:
from torch.ao.quantization import QuantStub, DeQuantStub class QuantizedResNet(nn.Module): def __init__(self, num_classes=10): super().__init__() self.backbone = models.resnet18(pretrained=True) self.backbone.fc = nn.Linear(512, num_classes) self.quant = QuantStub() self.dequant = DeQuantStub() def forward(self, x): x = self.quant(x) # 量化输入 x = self.backbone(x) x = self.dequant(x) # 反量化输出 return xQuantStub 会记录输入数据的范围,DeQuantStub 负责把最后的输出转回浮点用于计算 loss。注意:这两个节点本身不改变数据,它们只是标记“这里要插量化器”,真正起作用是在 prepare_qat 之后。
3.2 分步执行:fuse → prepare_qat → 微调 → convert
第一步:融合模块。对于 ResNet18,我们只需要融合基本的 ConvReLU2d 和 ConvBNReLU:
model = QuantizedResNet(num_classes=10).eval() # 融合 Conv + BN + ReLU(FX 模式可以自动找,Eager 模式需要手动指定) model.fuse_model() # 对 torchvision 自带结构有效如果你的模型不是标准结构,FX 模式更省心,它能够自动分析计算图并完成融合:
from torch.ao.quantization.quantize_fx import prepare_qat_fx # 先实例化原始模型 model = models.resnet18(pretrained=True) model.fc = nn.Linear(512, 10) # FX 模式自动融合 from torch.ao.quantization.fx.graph_module import fuse_fx model = fuse_fx(model) # 自动融合已知结构第二步:设置量化配置并执行 prepare_qat。量化配置决定了权重和激活用什么 observer、按张量还是按通道量化。我推荐这么设置:
import torch.ao.quantization as tq # 权重按通道量化(更精确),激活按张量量化(实际硬件更友好) qconfig = tq.QConfig( activation=tq.MinMaxObserver.with_args(dtype=torch.quint8, qscheme=torch.per_tensor_affine), weight=tq.MinMaxObserver.with_args(dtype=torch.qint8, qscheme=torch.per_channel_symmetric) ) model.qconfig = qconfig # Eager 模式 model = tq.prepare_qat(model, inplace=True) # FX 模式 # model = prepare_qat_fx(model, qconfig)这里要特别解释一下 qscheme 的选择:权重用 per_channel_symmetric,意思是每个输出通道有自己的 scale 和 zero_point,量化粒度更细,能显著减少误差;激活用 per_tensor_affine,因为绝大多数推理引擎(比如 ONNX Runtime、TensorRT)对激活只支持逐张量量化。在选型之前,务必要确认你的部署后端支持哪种方案,否则训练出来的“精度恢复”在转换后会被打回原形。
第三步:微调训练。QAT 不是重新训练,是在原有权重基础上做小幅调整。学习率千万别用大了,我一般用 1e-5 到 1e-6,训练 5~20 个 epoch 就够了。学习率太大,模型会直接偏离原始分布,精度不升反降。
optimizer = torch.optim.Adam(model.parameters(), lr=1e-5) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10) for epoch in range(10): model.train() for images, labels in train_loader: out = model(images) loss = nn.CrossEntropyLoss()(out, labels) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step()第四步:转换为真正的 INT8 模型。微调完之后,模型还是一个带伪量化节点的浮点模型,需要执行 convert 操作,把 FakeQuantize 的统计值固化成真正的 INT8 权重和 scale/zero_point:
# 先把模型切回 eval 模式,确保 BN 统计量固定 model.eval() # Eager 模式 quantized_model = tq.convert(model, inplace=False) # FX 模式 # from torch.ao.quantization.quantize_fx import convert_fx # quantized_model = convert_fx(model)convert 之后就得到一个 torch.ao.quantization.QuantizedModel,可以用 torch.jit 导出为 TorchScript 部署到生产环境。在导出之前一定要在验证集上对比三个东西:原始 FP32 精度、QAT 后未 convert 模型精度、convert 后 INT8 模型精度。多数情况下 QAT 后的 INT8 模型比 PTQ 的 INT8 模型要高好几个点,但比 FP32 还是略低一点,这很正常。
3.3 核心参数怎么调:observer 类型、批量大小、学习率
- observer 类型:MinMaxObserver 简单粗暴,直接用全局最小/最大值;MovingAverageMinMaxObserver 更平滑,对小批量数据震荡不敏感;HistogramObserver 用直方图估计百分位,适合有离群值的情况。我的经验:先用 MinMaxObserver 跑一遍,如果精度不达标再换 MovingAverageMinMaxObserver 或 HistogramObserver(用 99.99% 分位),不要一上来就上复杂的。
- 微调批量大小:尽量和训练时的 batch size 一致,至少不要小于 32,保证 BN 统计量稳定。
- 学习率策略:不推荐使用大学习率 + 预热(warmup)的组合,QAT 阶段模型已经足够收敛,再搞大规模预热等于把模型搅乱。线性衰减或余弦退火就够了。
4. 避坑实录:精度不达标最常见的 6 个原因与排查链路
QAT 不是“跑了就有效”,我自己在不同项目里踩过很多坑,这里总结出 6 个最容易让人栽跟头的地方,并按排查顺序列出来,你在实战中一旦发现精度不对,直接照着这个链路去查。
4.1 坑一:BN 层统计量“冻”了个寂寞
QAT 微调时,模型处于 train 模式,BN 层默认会持续更新 running_mean 和 running_var。这本来没问题,但如果 convert 之前没有把模型切到 eval 模式,BN 统计量还是半个“动态”状态,convert 时会用到最后一次迭代的统计量,可能和全局分布偏差很大,导致精度突然下跌。
正确做法:在 convert 前必须 model.eval(),让 BN 层使用固定的 running_mean / running_var。有些同学训练到最后一步忘记切换,结果 convert 后精度莫名其妙比 QAT 训练中的低很多,一排查才发现是这个原因。
4.2 坑二:模型结构太“花哨”,融合漏了关键模块
FX 模式能自动融合标准 Conv+BN+ReLU,但如果你在 forward 里写了自定义模块(比如自定义 pool、多分支相加、SE 模块),融合映射里根本没定义,它就“安静地跳过”了。漏融合的直接后果是量化点变多,误差增大。
排查方式:QAT 训练结束后,打印一下模型结构,数一数还有多少个独立的 BN 层——理论上融合完的模型不应该有独立 BN 层。如果发现还有 BN 层残留,你需要自己定义融合映射,或者干脆把这些层的量化跳过(在 qconfig 里设为 None),这比让量化器硬上更安全。
4.3 坑三:convert 出来的模型“数值不动”,模型好像失效了
“int8 量化后精度下降,数值不动”这个问题在边缘设备部署时特别多。我遇到过的原因主要有三种:
- observer 没更新:prepare_qat 之后如果模型一直在 eval 模式跑数据,observer 不会统计任何信息,统计出来的 scale 是默认值,convert 后的模型输出自然一片乱码。解决方法是确保 observer 在 train 模式下“见过”足够多的数据。
- MAX_VALUE 溢出:某些硬件对 INT8 的数值上限有额外限制,你在 PyTorch 里的 qconfig 是 [-128, 127],但转换到推理引擎时用的是 [-127, 127](去掉 -128),一旦权重落在这个边界附近,推理结果就会异常。发现这种情况需要重新校准,或者调整 observer 的 percentile。
- 导出格式问题:如果你用的是 custom 的推理引擎,convert 之后没有正确读取 zero_point,相当于你拿着 int8 的数值当成浮点去算,结果肯定是乱飘。
4.4 坑四:QAT 用的数据集太“干净”,过拟合到训练集上了
QAT 和普通微调一个最大的区别:它是在一个本来就已收敛的模型上做小幅调整,所以特别容易过拟合到微调数据集上。如果你的训练集和部署场景严重不一致,QAT 做过几个 epoch 之后,验证集精度反而比 PTQ 还低。
我的解法:QAT 的数据尽可能覆盖部署场景的分布,至少要包含校准集的数据。不要在 QAT 阶段使用严格的数据增强(随机裁剪、旋转都不要),因为增强后的分布会使模型重新适应“更花哨”的输入模式,偏离原始部署输入。
4.5 坑五:感知训练评估时用了 train 模式
QAT 训练中想快速看一眼当前精度,直接把模型切到 eval 模式去跑验证集,这是允许的;但如果你在训练还没有结束的时候,用 train 模式去评估,BN 是在动态变化的,FakeQuantize 也在实时更新统计量,精度看起来忽高忽低,很容易误导你做出“这么快就收敛了”的错误判断。
我的经验:QAT 训练过程中一定要用一个固定的 eval 评估流程,每次评估前都 dynamic 地做 model.eval(),评估完再切回 train。最后统计两组数据:eval 模式下的模型精度,convert 之后的 INT8 精度,两者差应该非常小。
4.6 坑六:量化敏感层全都在“硬抗”,导致整网精度被拖垮
不是每一层都适合量化。实际排查中我发现很多模型的第一层卷积(输入通常是 3 通道 RGB)和最后一层全连接(分类层)对量化异常敏感——前者是因为输入范围宽、变化极大;后者是因为输出 logits 直接决定分类结果,微小偏差都会被放大。
处理方式:在 qconfig 里把敏感层的 qconfig 设为 None,即保留 FP32(或后续用 FP16)。很多推理引擎允许混合精度部署,Quantizable 的层 + FP32 敏感层可以共存。这样做之后,量化带来的精度损失往往能再压缩一半以上。
5. 排查链路全流程:从“精度不对”到“找到真凶”
有时候问题不是一次性就能定位的,我建议你把下面这个排查链路打印出来贴在工位上,遇到 QAT 精度不对就照着走:
| 步骤 | 检查项 | 通过标准 |
|---|---|---|
| 1 | 原始 FP32 模型在验证集上的精度 | 和训练时一致,确保模型没退化 |
| 2 | PTQ 模型精度 | 记录基线,看 QAT 是否有提升空间 |
| 3 | QAT 微调结束(convert 前)evel 模式精度 | 应该高于 PTQ,且接近 FP32 |
| 4 | convert 后 INT8 模型精度 | 与第 3 步差 ≤1% |
| 5 | 导出到目标推理引擎后的精度 | 与第 4 步一致,若不一致,检查引擎的量化算子支持情况 |
| 6 | 部署后线上输入的精度 | 与第 5 步一致,若不正常,检查数据分布是否偏移 |
这条链路我反复用过很多次:先把问题定位到“哪一步开始精度掉”,再回到对应步骤去排查具体原因(融合、observer、BN 状态、敏感层等),基本能在 30 分钟内锁定根因。
6. 进阶调优与个人实践经验
如果 QAT 已经跑通但精度还差一口气,下面几个调优技巧可以直接拿来试。
6.1 先 PTQ 后 QAT:让量化参数有一个好的热身起点
很多开源代码里 QAT 都是从头训练或者从预训练权重开始。但更稳妥的做法是:先用少量校准集跑一遍 PTQ,把每一层的 scale 和 zero_point 统计好,然后把这些统计值作为 QAT 的初始值。这样 Observer 在 QAT 一开始就处于一个“见过世面”的状态,微调时能更快收敛。
PyTorch 实现这个流程的关键是:先 prepare_qat,然后手动跑几个 batch 的“校准数据”(只前向、不反向),让 observer 完成初始化统计,再进行正式微调和学习率调整。
6.2 用蒸馏辅助 QAT:让 FP32 老师带 INT8 学生
QAT 的训练目标除了交叉熵 loss 之外,可以额外加一个蒸馏 loss:拿原始 FP32 模型(老师)的输出和 QAT 模型(学生)的输出做 KL 散度对齐。这样做的好处是,学生不仅学习正确的标签,还学习老师对模糊样本的“软判断”,这比 hard label 下的精度恢复更稳。
具体实现只需要在微调的 loss 函数上做一点改动:
with torch.no_grad(): teacher_logits = teacher_model(images) student_logits = student_model(images) hard_loss = nn.CrossEntropyLoss()(student_logits, labels) soft_loss = nn.KLDivLoss(reduction="batchmean")( F.log_softmax(student_logits / T, dim=1), F.softmax(teacher_logits / T, dim=1) ) loss = hard_loss + alpha * (T * T) * soft_lossT 是蒸馏温度,通常取 3~8;alpha 是蒸馏损失的权重,我一般从 0.1 开始调。这个方法在轻量网络上效果尤其明显,MobileNet 上我见过把精度拉回 FP32 水平的情况。
6.3 分层决策:哪些层保留 FP32,哪些层必须量化
不同的推理硬件对“跳过量化”的支持程度不一样。在动手微调之前,先做个实验——把网络按层遍历,每次只量化其中一层,其他层保留 FP32,观测哪一层单独量化会导致精度大跌。把这几层标记为“敏感层”,要么跳过量化,要么在 QAT 训练时单独给更大的训练权重。
这种“逐层量化敏感度分析”看起来复杂,其实代码很少,几分钟就能跑完,但对最终部署精度的影响非常大。我之前在 YOLOv5 的项目里,就是用这个方法发现检测头的两三个卷积层是敏感层,跳过量化后 mAP 直接回升了 2 个多点。
6.4 评估 QAT 是否成功的“及格线”
我不建议只用 Top-1 Accuracy 单一指标评估 QAT 的好坏,尤其是回归类模型,精度掉了但数值分布可能无限接近——也有可能是评估方式的问题。建议从三个维度综合判断:
- 精度指标:Top-1 / Top-5 / mAP / IoU 等任务指标不能明显掉。
- 数值对齐度:对比 FP32 和 INT8 模型的输出向量,计算余弦相似度和平均绝对误差,一般余弦相似度 ≥0.99、MAE 极小才算达标。
- 层级 SQNR(信号量化噪声比):量化后每一层输出与原始 FP32 层输出的信噪比,若某层 SQNR 异常低,说明该层是敏感层,需要特殊处理。
把这三个维度记录下来,你的 QAT 结论就不再是“好像还行”,而是可量化、可追溯的。
7. 写在最后的操作心得
我自己的经验是:QAT 不是一个“银弹”,它不能 100% 保证 INT8 无损,但绝大多数情况下能把精度损失从“不能接受”压到“基本看不出来”。真正决定成败的往往不是那些花哨的技巧,而是最基础的几件事:有没有正确融合 Conv+BN+ReLU、学习率有没有过大、convert 前有没有切换 eval 模式、Observer 有没有在足够有代表性的数据上完成统计。
另外有一点想提醒大家:QAT 训练完模型,不要只盯着验证集精度,一定要尽早导出到目标推理引擎跑一遍端到端验证。因为 PyTorch 里 convert 出来的量化模型和最终硬件上执行的算子可能不完全等价,早发现早处理,别等到部署到设备上了再回来排查。
最后再分享一个小技巧:把 QAT 微调过程中的模型检查点(checkpoint)按 epoch 保存下来,不要只保留最后一个。因为 QAT 微调后期很容易出现过拟合或精度波动,有一个中间的检查点精度反而更高。我一般每 2 个 epoch 保存一次 checkpoint,微调结束后统一在验证集上评估,选最优的那个去 convert。这个方法几乎零成本,但经常能让你多挽回 0.5~1 个点的精度。
如果这篇文章能帮你把模型量化后那口“恶气”吐出来,我就很满足了。有任何 QAT 相关的问题,欢迎在评论区交流,我看到会尽量回复。