ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch量化感知训练实战:让INT8模型精度不再掉点

PyTorch量化感知训练实战:让INT8模型精度不再掉点 这一篇我想跟各位认真聊聊量化感知训练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 掉点情况原因大 CNNResNet50/101 等通常 0.5%~3%层数深但结构规整误差相对可控轻量 CNNMobileNet、ShuffleNet3%~8%甚至更多深度可分离卷积对量化极敏感数值范围碎片化TransformerBERT、ViT2%~10%Softmax、LayerNorm、残差结构容易放大误差检测/分割模型YOLO、DeepLab2%~8%多任务头、特征金字塔导致逐层误差累加复杂RNN/LSTM很容易崩循环结构内误差反复叠加时序依赖敏感如果你手头的模型 PTQ 只掉 0.5 个点说实话没必要上 QAT用一些校准技巧就能救回来但如果你掉 3 个点以上尤其是轻量网络或者 Transformer老老实实走 QAT 是性价比最高的方案。2. QAT 为什么能拯救精度从“被动接受误差”到“主动适应误差”PTQ 是模型已经训练完了再拿一批数据去统计量化参数模型本身对量化这件事一无所知。QAT 的思路反过来让模型在训练的时候就“提前体验”量化带来的误差迫使权重适应一个更粗糙的数值环境。用一个不太严谨但很好理解的类比——QAT 就像让一个习惯看高清水牌的人先戴一副磨砂眼镜去背招牌等他摘掉眼镜的时候反而能更准确地认出低分辨率下的字。2.1 伪量化节点FakeQuantize到底做了什么QAT 的核心是 FakeQuantize 模块。它的作用是在前向传播的时候把浮点数值模拟成“量化后再反量化”的结果——先量化到整数模拟精度损失再反量化回浮点保证后续计算还是浮点。这样一来模型在训练时的前向计算就带上了真实的量化噪声。关键是反向传播。量化函数本身不可导它是个阶梯函数几乎处处导数为 0所以 PyTorch 里用了直通估计器STEStraight-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结尾加 DeQuantStubfrom torch.ao.quantization import QuantStub, DeQuantStub class QuantizedResNet(nn.Module): def __init__(self, num_classes10): super().__init__() self.backbone models.resnet18(pretrainedTrue) 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 和 ConvBNReLUmodel QuantizedResNet(num_classes10).eval() # 融合 Conv BN ReLUFX 模式可以自动找Eager 模式需要手动指定 model.fuse_model() # 对 torchvision 自带结构有效如果你的模型不是标准结构FX 模式更省心它能够自动分析计算图并完成融合from torch.ao.quantization.quantize_fx import prepare_qat_fx # 先实例化原始模型 model models.resnet18(pretrainedTrue) 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( activationtq.MinMaxObserver.with_args(dtypetorch.quint8, qschemetorch.per_tensor_affine), weighttq.MinMaxObserver.with_args(dtypetorch.qint8, qschemetorch.per_channel_symmetric) ) model.qconfig qconfig # Eager 模式 model tq.prepare_qat(model, inplaceTrue) # 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(), lr1e-5) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max10) 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, inplaceFalse) # 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 模式能自动融合标准 ConvBNReLU但如果你在 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 模型在验证集上的精度和训练时一致确保模型没退化2PTQ 模型精度记录基线看 QAT 是否有提升空间3QAT 微调结束convert 前evel 模式精度应该高于 PTQ且接近 FP324convert 后 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(reductionbatchmean)( F.log_softmax(student_logits / T, dim1), F.softmax(teacher_logits / T, dim1) ) loss hard_loss alpha * (T * T) * soft_lossT 是蒸馏温度通常取 3~8alpha 是蒸馏损失的权重我一般从 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 无损但绝大多数情况下能把精度损失从“不能接受”压到“基本看不出来”。真正决定成败的往往不是那些花哨的技巧而是最基础的几件事有没有正确融合 ConvBNReLU、学习率有没有过大、convert 前有没有切换 eval 模式、Observer 有没有在足够有代表性的数据上完成统计。另外有一点想提醒大家QAT 训练完模型不要只盯着验证集精度一定要尽早导出到目标推理引擎跑一遍端到端验证。因为 PyTorch 里 convert 出来的量化模型和最终硬件上执行的算子可能不完全等价早发现早处理别等到部署到设备上了再回来排查。最后再分享一个小技巧把 QAT 微调过程中的模型检查点checkpoint按 epoch 保存下来不要只保留最后一个。因为 QAT 微调后期很容易出现过拟合或精度波动有一个中间的检查点精度反而更高。我一般每 2 个 epoch 保存一次 checkpoint微调结束后统一在验证集上评估选最优的那个去 convert。这个方法几乎零成本但经常能让你多挽回 0.5~1 个点的精度。如果这篇文章能帮你把模型量化后那口“恶气”吐出来我就很满足了。有任何 QAT 相关的问题欢迎在评论区交流我看到会尽量回复。
返回列表