量化感知训练QAT这件事我前前后后在三四个项目里踩过坑从最早把torch.quantization当成黑盒用到后来被精度掉点折磨得怀疑人生再到现在能比较从容地判断这个模型该不该上QAT、该在哪个位置插fake quant。这篇就把我积累的这些经验完整摊开讲一遍尽量做到你看完能直接上手改自己的模型。先说清楚QAT到底解决什么问题。PyTorch的量化分两条路训练后量化PTQ和量化感知训练QAT。PTQ就是拿一个训好的浮点模型直接校准一下统计量然后转成int8快是快但遇到depthwise卷积多、激活分布尖锐、或者模型本身比较小的场景精度掉得很难看。QAT的思路是在训练阶段就模拟量化的误差让网络权重去适应这种误差最后推理时用真正的int8算子精度通常能拉回PTQ掉的大部分甚至全部。适合谁看已经会训PyTorch模型、想往端侧部署走、但被PTQ精度劝退的人。1. 先搞清楚QAT在PyTorch里到底动了哪些手脚1.1 量化的数学本质仿射映射与量化参数要理解QAT得先接受一个事实int8量化本质是一个仿射映射。把一个浮点张量x映射到int8公式是x_int round(x / scale) zero_point x_dequant (x_int - zero_point) * scalescale是缩放因子zero_point是零点偏移。这两个东西合起来叫量化参数quantization parameters。scale决定了量化步长zero_point保证浮点里的0能精确映射到某个整数上——这点很关键因为padding用的就是0如果0映射不准卷积边界会引入系统性误差。PyTorch里默认用的是per-tensor还是per-channel权重量化默认是per-channel每个输出通道一组scale/zero_point激活量化默认是per-tensor。为什么这么设计因为权重的分布在不同通道之间差异很大per-channel能显著降低量化误差而激活在推理时是动态的per-tensor实现起来更高效硬件也更友好。这个默认值你在做QAT时基本不用改但要知道它的存在否则调精度时会一头雾水。1.2 fake quantQAT的核心机关QAT不是真的在训练时用int8算而是在浮点计算图里插入fake quant模块。前向的时候fake quant做的是量化再反量化——把浮点值按上面的公式压到int8的格点上再还原回浮点。这样前向的输出就带上了量化误差反向传播时因为round的梯度几乎处处为0PyTorch用的是STEStraight-Through Estimator也就是把fake quant当成恒等映射来传梯度。你可以把fake quant理解成一个带噪声的恒等层数值上它把值往格点上吸梯度上它假装什么都没发生。网络在训练中慢慢学会把权重和激活调整到那些量化后也不会太失真的位置。这就是QAT能救回精度的根本原因。1.3 observer统计量是怎么攒出来的fake quant要工作得知道scale和zero_point。这两个值由observer负责统计。PyTorch里常见的observer有MinMaxObserver、MovingAverageMinMaxObserver、HistogramObserver等。QAT阶段默认用滑动平均的min-max observer因为它能在训练过程中持续更新统计量而不是像PTQ那样一次性校准。这里有个容易忽略的点observer统计的是每一批数据的min/max然后用滑动平均平滑。所以QAT训练时的数据分布必须和真实推理分布接近否则统计出来的scale会偏。我见过有人拿训练集的一个小子集跑QAT结果部署时精度崩了——因为那个子集的激活范围跟真实场景差太远。2. 从浮点模型到QAT模型配置流程的每一步为什么这么写2.1 环境与版本别在这上面栽跟头PyTorch的量化API在不同版本之间改动很大尤其是torch.ao.quantization取代torch.quantization之后。我的建议是锁定一个版本别追新。目前比较稳的是1.13到2.1这个区间torch.ao.quantization的接口已经比较成熟。安装就正常装pip install torch2.1.0 torchvision0.16.0CPU推理的量化不需要额外依赖但如果要跑在特定后端上得确认后端支持的算子集。比如x86上常用fbgemmARM上常用qnnpack。这两个后端对算子的支持不完全一样选错了会在convert阶段报operator not supported。提示fbgemm和qnnpack的选择不是随便挑的。x86服务器选fbgemmARM移动端选qnnpack这是官方推荐也是实测最稳的组合。2.2 模型准备哪些层能量化哪些不能不是所有层都能量化。PyTorch的量化只支持特定模块nn.Conv2d、nn.Linear、nn.ReLU、nn.BatchNorm2d会被折叠、nn.MaxPool2d等。像nn.LSTM、自定义的复杂算子、某些attention实现默认是不支持的。所以第一步是改造模型结构把能融合的算子融合掉。最常见的是ConvBNReLU的融合import torch from torch.ao.quantization import fuse_modules model.eval() model_fused fuse_modules( model, [[conv1, bn1, relu1], [conv2, bn2, relu2]], inplaceFalse )为什么要融合因为BN在推理时本质是一个逐通道的仿射变换可以完全吸收进前面的卷积权重里。融合之后量化只需要处理一个Conv误差来源少了一个精度更稳推理也更快。这一步在PTQ和QAT里都是必须的但很多人做PTQ时忘了融合精度掉了还找不到原因。2.3 qconfig把量化方案声明出来qconfig是QAT配置的核心它规定了权重和激活分别用什么observer、什么量化方案。一个典型的QAT qconfig长这样from torch.ao.quantization import QConfig, FakeQuantize from torch.ao.quantization.observer import MovingAverageMinMaxObserver qconfig QConfig( activationFakeQuantize.with_args( observerMovingAverageMinMaxObserver, quant_min0, quant_max255, dtypetorch.quint8, qschemetorch.per_tensor_affine, reduce_rangeFalse ), weightFakeQuantize.with_args( observerMovingAverageMinMaxObserver, quant_min-128, quant_max127, dtypetorch.qint8, qschemetorch.per_channel_symmetric, reduce_rangeFalse ) )注意激活用的是quint8无符号0到255权重用的是qint8有符号-128到127。为什么激活用无符号因为ReLU之后的激活都是非负的用无符号能多利用一位精度。权重有正有负所以用有符号。reduce_range这个参数值得说一句。在某些老硬件上int8乘法的中间结果会溢出所以要把范围砍一半激活变成0到127权重变成-64到63。现代硬件基本不需要但如果你部署到比较老的设备上可能要打开它。打开之后精度会掉一点这是代价。2.4 prepare插入fake quant进入模拟量化状态配置好qconfig之后调用preparefrom torch.ao.quantization import prepare_qat, convert model_fused.qconfig qconfig model_qat prepare_qat(model_fused, inplaceFalse)prepare_qat做的事情是遍历模型给每个支持量化的模块插入fake quant模块并挂上observer。这时候模型的前向计算已经带上了量化模拟但权重还是浮点的。有个细节prepare_qat之后模型里会多出很多fake_quant子模块state_dict的key也会变。如果你要从一个已经训好的浮点模型加载权重必须在prepare_qat之前加载或者用load_state_dict时设置strictFalse并手动对齐key。我一般是在融合之后、prepare之前把浮点权重load进去这样最干净。3. QAT训练阶段学习率、epoch数和那些反直觉的调参经验3.1 训练策略为什么QAT不需要训太久QAT的本质是微调不是从头训。因为浮点模型已经学到了好的特征表示QAT只是让权重去适应量化误差。所以通常只需要原训练epoch的10%左右甚至更少。我的经验值图像分类任务原训练100 epoch的模型QAT跑5到10个epoch就够了。检测和分割任务因为对定位精度敏感可能要15到20个epoch。跑太多反而可能过拟合因为fake quant引入的噪声会让模型在训练集上表现变差。学习率方面必须用比原训练小得多的学习率。我一般用原训练初始学习率的1/100到1/10。比如原训练用0.1QAT就用0.001到0.01。为什么因为QAT是在一个已经很优的解附近做微调学习率太大会把权重踢出好的区域精度反而下降。optimizer torch.optim.SGD( model_qat.parameters(), lr1e-4, # 比原训练小两个数量级 momentum0.9, weight_decay1e-5 ) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max10 )3.2 冻结observer什么时候该停止更新统计量QAT训练到后期有一个关键操作冻结observer。也就是让scale和zero_point不再更新固定下来。为什么要冻结因为如果统计量一直在变模型永远在追一个移动的目标很难收敛到稳定状态。PyTorch提供了freeze_observerfrom torch.ao.quantization import freeze_observer # 训练到第7个epoch时冻结 if epoch 7: model_qat.apply(freeze_observer)冻结的时机怎么定我的做法是总epoch的70%到80%处冻结。比如跑10个epoch第7或第8个epoch冻结。冻结之后再跑2到3个epoch让模型适应固定的量化参数。这个节奏实测比较稳。注意冻结observer之后学习率最好也降一档因为此时量化参数固定了模型需要更精细的调整。3.3 一个反直觉的点QAT训练时精度会掉很多人第一次做QAT会慌怎么训练集精度比浮点模型低了好几个点这是正常的。因为fake quant在前向引入了量化误差训练时的精度本身就带了这个误差。你要看的是convert之后的int8模型精度而不是QAT训练过程中的浮点精度。我一般会在训练过程中定期做一次模拟convert来监控把模型转成eval模式跑一遍验证集看精度趋势。如果QAT训练精度稳定在一个比浮点低1到2个点的水平convert之后通常能回到接近浮点的水平。如果QAT训练精度一路往下掉那说明学习率太大或者数据有问题。4. convert与部署验证从模拟到真实的最后一公里4.1 convert把fake quant换成真算子训练完成后先切到eval模式然后convertmodel_qat.eval() model_int8 convert(model_qat, inplaceFalse)convert做的事情是把fake quant模块替换成真正的量化算子权重从浮点转成int8激活的量化参数固化到算子内部。转换后的模型是真正的int8模型推理时用的是量化kernel。这里有个坑convert必须在eval模式下做。如果模型还在train模式BN的统计量不对convert出来的结果会错。而且convert之后模型不能再训练了所以顺序不能乱。4.2 精度验证怎么判断QAT成功了convert之后跑一遍验证集和浮点模型对比。我的验收标准是指标可接受范围理想范围精度掉点 1% 0.3%模型大小约为浮点的1/41/4推理延迟有明显下降下降2到4倍如果掉点超过1%先别急着放弃按下面的顺序排查检查是否有层没被量化用model_int8打印结构看哪些还是浮点检查observer冻结时机是否太早检查qconfig的reduce_range是否被误开尝试per-channel激活量化部分后端支持4.3 部署时的算子兼容性检查convert成功不代表能部署。目标后端必须支持模型里所有的量化算子。我一般用torch.ao.quantization.quantize_fx的backend_config来指定后端然后在convert时它会做算子检查。如果报某个算子不支持有两个选择一是把这个算子排除在量化之外用set_observed_module或者自定义qconfig二是换一个支持该算子的后端。前者会损失一点性能但能保证跑通后者要看硬件条件。提示depthwise卷积在很多后端上的量化支持都不太好如果你的模型里有大量depthwise比如MobileNet系列要特别关注这一块的精度和兼容性。5. 那些文档里不会写的踩坑记录5.1 坑一BatchNorm在QAT里的微妙行为BN在QAT里是个麻烦制造者。前面说了ConvBN要融合但如果你在QAT训练时BN还在更新running stats融合就会出问题。我的做法是融合之后把BN的running stats固定住或者在QAT训练时用很小的momentum。更稳妥的做法是融合后直接删掉BN层因为它的参数已经吸收进Conv了。但PyTorch的fuse_modules默认是保留BN结构的只是把计算合并所以convert时它知道怎么处理。如果你手动删了BN反而可能让convert找不到对应的模式。5.2 坑二自定义模块的量化处理模型里如果有自定义模块默认是不会被量化的。你需要给它写一个from_float类方法告诉PyTorch怎么把它转成量化版本。这个工作量大而且容易出错。我的建议是能改结构就改结构把自定义模块拆成PyTorch原生支持的算子组合。实在改不了的就把它排除在量化之外接受这部分用浮点计算。混合精度推理在很多后端上是支持的只是性能提升会打折扣。5.3 坑三数据加载器的一致性QAT训练用的数据增强必须和浮点训练保持一致尤其是归一化参数。因为observer统计的是归一化之后的激活范围如果归一化参数变了scale就全错了。我踩过一次浮点训练用ImageNet的mean/stdQAT时图省事用了另一套参数结果convert后精度掉了5个点。排查了半天才发现是归一化的问题。所以数据预处理这块直接复用浮点训练的配置别改。5.4 坑四多卡训练与observer同步如果用DataParallel或多卡训练observer的统计量在不同卡上是独立的可能导致scale不一致。PyTorch的QAT对多卡支持有限我的经验是尽量单卡做QAT或者用DistributedDataParallel并手动同步observer的统计量。单卡虽然慢但省心。6. 进阶把QAT嵌进自己的训练框架6.1 用FX Graph Mode做更细粒度的控制前面讲的都是Eager Mode的QAT它的缺点是只能量化预定义的模块对动态控制流支持差。FX Graph Mode通过追踪计算图来做量化能处理更复杂的模型结构。from torch.ao.quantization.quantize_fx import prepare_qat_fx, convert_fx qconfig_dict { : qconfig, module_name: [ (attention, None), # attention模块不量化 ] } model_prepared prepare_qat_fx(model_fused, qconfig_dict) # ... 训练 ... model_int8 convert_fx(model_prepared)qconfig_dict里的None表示该模块不量化这给了很细的控制粒度。FX模式对transformer类模型的量化支持比Eager模式好很多如果你在做NLP模型的量化建议直接上FX。6.2 混合精度量化不是所有层都值得int8有些层对量化特别敏感比如第一层卷积和最后一层全连接。第一层直接接触输入量化误差会被放大最后一层直接决定输出量化误差直接影响结果。这两层有时候保持浮点反而更好。在FX模式下可以这样配置qconfig_dict { : qconfig, module_name: [ (conv1, None), # 第一层不量化 (fc, None), # 最后一层不量化 ] }代价是这两层还是浮点计算推理时会插入dequant/quant转换有一点额外开销。但如果能救回精度这点开销值得。6.3 监控量化误差一个实用的小工具我写了个小函数用来对比每一层量化前后的输出差异定位是哪一层引入了大误差def compare_layer_outputs(float_model, quant_model, input_tensor): float_model.eval() quant_model.eval() float_outputs {} quant_outputs {} def get_hook(name, storage): def hook(module, input, output): storage[name] output.detach() return hook for name, module in float_model.named_modules(): if isinstance(module, (torch.nn.Conv2d, torch.nn.Linear)): module.register_forward_hook(get_hook(name, float_outputs)) for name, module in quant_model.named_modules(): if isinstance(module, (torch.nn.Conv2d, torch.nn.Linear)): module.register_forward_hook(get_hook(name, quant_outputs)) with torch.no_grad(): float_model(input_tensor) quant_model(input_tensor) for name in float_outputs: if name in quant_outputs: diff (float_outputs[name] - quant_outputs[name]).abs().mean() print(f{name}: mean abs diff {diff:.6f})跑一遍就能看出哪层误差最大。如果某一层的误差明显高于其他层就考虑把它排除在量化之外或者调整它的observer配置。7. 关于QAT值不值得做的判断最后说点实在的。QAT不是万能的也不是所有场景都值得做。我的判断标准是如果PTQ掉点在1个点以内直接上PTQ别折腾QAT。QAT要改训练流程、要调参、要重新训时间成本不低。只有当PTQ掉点超过2个点或者模型对精度极其敏感比如医疗影像、自动驾驶感知才值得上QAT。另外QAT的效果和模型结构强相关。大模型ResNet50以上QAT后基本能追平浮点小模型MobileNet级别因为本身冗余度低QAT后可能还是掉0.5到1个点这是结构决定的不是调参能解决的。我在实际项目里的体会是QAT最大的价值不是把精度拉满而是把精度拉到一个可接受的稳定水平。它给你的是一个可控的精度-性能权衡点而不是免费的午餐。理解这一点你在做量化决策时就不会那么焦虑了。