简介:面向需要部署视觉Transformer的工程师与算法研究者,这个量化加速工程包围绕ViT、DeiT、SwinT三种模型的后训练量化(PTQ)展开,解决模型在资源受限环境下推理耗时与内存占用过高的问题。资源共15个文件、压缩包仅41KB,以14个Python脚本和1个Markdown说明为主;Python代码覆盖模型定义、量化校准、整数化处理、量化层实现(卷积、线性、矩阵乘)、数据集封装及多个测试入口,Markdown文档则说明量化流程与复现步骤。PTQ无需重新训练模型,配合少量校准数据即可完成从浮点权重到低比特整数的转换,适合边缘设备与服务端推理加速场景。示例目录提供多组测试脚本,可分别验证单模型效果、整体流程及消融对比,方便评估量化前后精度与速度变化。包内目录结构清晰,按配置、量化层、工具函数、示例测试等模块组织,便于二次开发与实验对比;目前已有196人学习或下载,对算法工程师、部署工程师和相关方向研究生均有参考价值。
1. 量化加速ViT家族,先解决“FP32跑不动”的部署问题
把ViT从FP32压到INT8,听起来就是把权重和激活乘个scale再取整,但真正在ViT/DeiT/SwinT上跑一遍PTQ量化加速,就会发现最难的不是量化算法,而是模型结构里的LayerNorm、GELU、attention softmax一个接一个地在精度和速度之间出难题。下面按落地顺序讲清楚:从ViT结构里哪几个算子必须单独处理,到用ONNX Runtime静态量化跑通ViT、DeiT、SwinT三个模型,再到精度掉点时的排查顺序和参数调整。适合做端侧或CPU推理部署、手里有预训练ViT系列模型但不想走QAT重训的工程师。全程只依赖timm和onnxruntime,拿自己的分类任务数据就能复现。项目源码里用到的导出、校准、量化、验证四段代码,我会全部贴出来,照着顺序跑就能看到从FP32到INT8的完整变化。
2. 看懂PTQ量化边界:ViT/DeiT/SwinT到底哪里难量化
2.1 线性量化坐标:Scale、ZeroPoint与per-channel在哪生效
PTQ量化加速的第一步不是跑代码,而是搞清楚量化前后数值怎么对应。8bit线性量化的核心是round操作:量化值q = clamp(round(x / scale) + zero_point),反量化x = (q - zero_point) * scale。其中scale是对浮点数值range的等分粒度,zero_point是float零点映射到整数的偏移。模型推理时把FP32权重和激活换成INT8,计算完再反量化为浮点,中间多出的误差就是量化误差。
我用一段极简代码先把这个映射固定住,后面校准配置全都围绕它展开:
def quantize_tensor(x, scale, zero_point, bits=8): qmin, qmax = 0, (1 << bits) - 1 q = torch.round(x / scale + zero_point).clamp(qmin, qmax) return q.to(torch.int8 if bits == 8 else torch.int16) def dequantize_tensor(q, scale, zero_point): return (q.float() - zero_point) * scale这段代码逻辑很简单:除以scale再加zero_point,然后clamp到整数取值范围。注意zero_point在非对称量化里是个float,在对称量化里固定为0;ViT的权重一般用对称量化(zero_point=0),激活因为分布不一定过零点,用非对称量化更合适。per-channel和per-tensor的差异在于scale是按整张表共享一份,还是按输出通道各算一份:self-attention里的Linear weight是[out_dim, in_dim],per-channel能更好贴合不同通道的数值范围,尤其对ViT这种embedding维度很高的结构,per-channel基本是必选项。
ONNX Runtime的量化器在生成QDQ格式模型时,会把scale和zero_point作为常量节点插入到Dequantize算子旁边。也就是说,你不需要手工去填每一层的scale,量化器会根据校准阶段观察到的激活分布和权重分布自动计算。你需要操心的只是选择哪种校准方法和量化粒度,这两个参数决定了scale最终落在哪个区间。
2.2 ViT结构里的三个量化黑洞:LayerNorm、GELU、Softmax
ViT主流技术路线里,一个标准的ViT Block包含LayerNorm、MHA和MLP,中间穿插GELU。这三个算子恰好是量化的三个难点,几乎每次部署都要为它们单独做配置。
LayerNorm是对每个token的整条特征向量做归一化,输出被拉到一个动态范围很小的区间。初看数值范围稳定,但ViT不同层LayerNorm的输出均值差异很大,如果全部用同一个scale去量化,浅层和深层之间会出现系统性偏差。更关键的是LayerNorm内部有除法、平方根、减法,ONNX里的LayerNormalization算子如果在INT8下计算,误差会被后面attention放大。我一般会在量化配置里默认让它回退FP32,而不是硬量化。
GELU在负半轴有一段软饱和区,x<0时输出不是完全为0,而是趋近0。这意味着负半轴的细微数值变化在做round后容易全部变成0,导致激活信息丢失。ViT又是重度依赖非线性传播的结构,GELU如果被MinMax校准的极端值带偏,精度掉1~2个百分点很常见。处理方式是校准方法从MinMax换成MSE或Percentile,让裁剪点离开极端值区域,保留下半轴的细节信息。
Softmax在attention内部,输出0到1之间,看起来是最好量化的区间,但实际是坑。attention score在softmax之前的数值分布非常不均匀,有的接近0,有的接近负几十,全图量化后softmax输出被钝化,注意力权重容易被抹平。现在主流的做法是softmax保留FP32,只量化它前后的matmul,这也是我导出模型时设置opset17、用QDQ格式的原因:QDQ格式允许算子级别的混合精度回退,量化器在遇到不支持的算子时可以自动跳过。
这三个操作叠加起来,决定了ViT家族不能照搬CNN项目里那种“全局量化+看精度”的流程。CNN里的BatchNorm可以很自然地融合进卷积量化,ReLU的截断性质也让激活分布天然偏向0~正区间;ViT没有这些便利,LayerNorm、GELU、Softmax都需要单独看一遍再决定量不量化。
2.3 DeiT和SwinT额外添乱的细节:distill token与window shift
DeiT在ViT基础上加了distillation token,forward里除了classifier还会走distill head。导出ONNX时如果不处理,会得到多一个输出的模型,推理阶段还要额外算一个head。常见的做法是导出前只保留分类头,或者导出后从graph里删掉distill相关分支,否则量化校准阶段不仅多跑一遍计算,还可能把distill分支的数值范围混进校准统计,污染主分类头的scale。
SwinT则是给了另一套麻烦。它用窗口注意力把7x7的字块划成一个个窗口,再通过shift window跨层交互。这导致中间activation的shape和数值范围随着stage切换频繁变化,量化校准阶段如果只用全局scale,不同stage之间的动态范围冲突会让整体精度掉一截。我在实际项目里的处理是:对SwinT的量化先把校准数据量拉到500张以上,再用MSE校准,最后如果精度还是不高,就显式把window attention内部的matmul保留FP32。
这三个结构特点叠加在一起,也解释了为什么不能拿CNN那套“直接量化、看两眼精度”的经验直接用在ViT家族上。第3、4章把这一套边界落到导出和量化配置里的过程,一步步拆开看。
3. 把模型导出成PTQ可用的ONNX:从PyTorch到ONNX的一路排查
3.1 三条PTQ路线怎么选:PyTorch FX、ONNX Runtime与TensorRT
对ViT做PTQ量化加速,工具链上有三条主流路线,很多人一开始就卡在“到底用哪个”上。PyTorch FX量化走的是torch.ao.quantization的FX图模式,在x86 CPU上推理比较顺手,但ViT里LayerNorm、Softmax、GELU这几个算子经常需要手动加入白名单,否则trace阶段就报错,调起来比较费劲,适合模型最终要跑在PyTorch环境里的场景。TensorRT INT8效果最好,但依赖NVIDIA GPU和TensorRT版本,Calibrator的校准流程和onnx导出有版本耦合,一般放到最后再考虑。
ONNX Runtime静态量化是目前覆盖ViT/DeiT/SwinT成本最低的一条路:先把PyTorch模型导出成ONNX,再用onnxruntime.quantization做静态PTQ。好处是算子覆盖广、QDQ格式天然支持算子级回退,CPU/GPU/Mobile都吃同一份产物,也是现在ViT主流技术路线里部署侧用得最多的方案。就普通分类项目而言,用这条路线能一步到位,且不用重训、不用动训练代码。如果PTQ精度损失超过5个点再考虑上QAT,PTQ阶段能把浮点和INT8差值控制在2个点以内,通常没必要重训。
3.2 用timm导出ViT/DeiT/SwinT的ONNX并跑通推理
我用timm加载预训练模型,统一导出成ONNX。三个模型的加载方式一致,只在create_model的参数上区分:
import timm import torch def load_and_export(model_name, onnx_path): # exportable=True 会把timm内部不稳定的自定义op替换成标准onnx算子 model = timm.create_model(model_name, pretrained=True, exportable=True).eval() dummy = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy, onnx_path, input_names=['input'], output_names=['logits'], dynamic_axes={'input': {0: 'batch'}, 'logits': {0: 'batch'}}, opset_version=17, # QDQ格式量化需要opset>=17 ) load_and_export('vit_base_patch16_224', 'vit_base_224.onnx') load_and_export('deit_base_patch16_224', 'deit_base_224.onnx') load_and_export('swin_base_patch4_window7_224', 'swin_base_224.onnx')这里有两个关键参数。exportable=True:timm在导出时会替换掉fx trace不稳定的自定义算子,把部分inplace操作和attention的实现改成标准onnx算子,这一步能挡掉大部分“Could not export Python operator”的报错。opset_version=17:ONNXRuntime静态量化在QDQ格式下需要17以上的opset,层叠的LayerNormalization和Softmax才能被后续工具识别并按节点回退。dynamic_axes让batch维度可变,校准阶段一次喂多张图,推理阶段一次喂一张,不用维护两个模型文件。
导出后先用ONNXRuntime跑一遍作为基线,这一步非常关键:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession('deit_base_224.onnx', providers=['CPUExecutionProvider']) inp = np.random.randn(1, 3, 224, 224).astype(np.float32) logits = sess.run(None, {'input': inp})[0] print('onnx output shape:', logits.shape)如果这里能正常输出,FP32基线就立住了。后面量化精度对比都以这个ONNX推理结果作为参考,不再回到PyTorch模型上做对比。这一步是很多项目源码里容易被忽略的环节:直接拿PyTorch的eval精度和ONNX INT8的精度对比,中间差了FP32 ONNX本身的转换损耗,会让排查精度问题时找错方向。我自己吃过这个亏,后来固定先确认FP32 ONNX的精度和PyTorch原始模型误差在0.1%以内,再继续往后走。
4. 用ONNX Runtime对ViT做静态PTQ:校准与量化全流程
4.1 校准数据怎么准备:300张无标注图片就够用
静态PTQ的关键输入不是标注,而是校准数据。校准数据的用途是统计每一层激活的数值分布,从而确定scale和zero_point,因此不需要分类标签,只要图片内容分布接近真实部署场景。以ImageNet训练的模型为例,收集300~500张来自不同类别、不同场景、不同光照的图片就能把激活分布统计得差不多;如果只有100张甚至几十张,attention的matmul层scale偏差会明显变大,SwinT这类多stage模型尤其敏感。
我自己一般会从训练集或测试集的子集里随机抽,避免全部来自同一摄像机或同一批拍摄环境。代码上,为了不依赖ImageFolder的目录结构,我直接写一个读图片路径的Dataset:
from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T TF = T.Compose([ T.Resize(256), T.CenterCrop(224), T.ToTensor(), T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) class CalibDataset(Dataset): def __init__(self, paths): self.paths = paths def __len__(self): return len(self.paths) def __getitem__(self, idx): img = Image.open(self.paths[idx]).convert('RGB') return TF(img), 0 # 300~500张图的路径列表,类别分布尽量均匀 paths = [...] loader = DataLoader(CalibDataset(paths), batch_size=8, shuffle=True, num_workers=4)transform必须和模型训练时一致,ViT家族用的都是224x224,均值和标准差也用ImageNet的默认值。batch_size建议8~16,太小会让每次校准的统计波动大,太大则内存占用高且对校准精度提升有限。校准数据是一次性消费的,量化器跑完一遍后读取器返回None即可,不需要反复重放。
4.2 实现CalibrationDataReader并喂给量化器
ONNX Runtime的静态量化需要数据读取器,接口只有一个get_next(),每次返回一个字典,key是ONNX模型的输入名,value是numpy数组。注意输入名要和导出时定义的input一致,否则量化器在读数据时会报key mismatch。
from onnxruntime.quantization import CalibrationDataReader class ViTCalibReader(CalibrationDataReader): def __init__(self, loader): self.loader = loader self.iter = iter(loader) def get_next(self): try: batch, _ = next(self.iter) return {'input': batch.numpy().astype(np.float32)} except StopIteration: return None calib_reader = ViTCalibReader(loader)注意这里有个细节,batch.numpy()出来的numpy数组是C-contiguous的,ONNX Runtime运行时要求输入内存连续;如果从Dataset里经过ToTensor后直接转numpy,这个条件天然满足。如果中间加了别的变换导致非连续内存,要先np.ascontiguousarray再返回,否则量化或推理阶段可能触发sigmoid fault。这个CalibrationDataReader会在量化器内部被循环调用,直到返回None为止,所以校准集用完一遍后迭代器自然耗尽即可。
4.3 量化配置逐项解释:MSE、per-channel、QDQ与MatMulConstBOnly
接下来是量化主流程。onnxruntime.quantization的quantize_static有非常多的参数,但ViT家族实际需要关注的就是下面这几个:
from onnxruntime.quantization import ( quantize_static, QuantType, CalibrationMethod, QuantFormat ) quantize_static( model_input='deit_base_224.onnx', model_output='deit_base_224_int8.onnx', calibration_data_reader=calib_reader, quant_format=QuantFormat.QDQ, per_channel=True, activation_type=QuantType.QUInt8, weight_type=QuantType.QInt8, calibrate_method=CalibrationMethod.MSE, extra_options={ 'ActivationSymmetric': False, 'WeightSymmetric': True, 'MatMulConstBOnly': True, 'QDQOpTypeFallback': True, }, )逐个说明。quant_format=QDQ:Quantize-Dequantize格式,量化后的模型里每个量化算子前后显式插入Dequantize节点,这样算子级回退、混合精度和后续的优化都更灵活。per_channel=True:权重按输出通道各一份scale,ViT的Linear层维度高,这一步不打开精度损失会明显。activation_type=QUInt8:激活用非对称无符号8bit,因为经过Normalize和LayerNorm之后激活分布不是对称的;weight_type=QInt8:权重用对称有符号8bit,这是CNN时代留下的经验,也适用于ViT的Linear和Conv。calibrate_method=CalibrationMethod.MSE:比MinMax对长尾分布更鲁棒,ViT的激活在attention score附近有明显长尾,MinMax会把动态范围拉宽,MSE能在整体重构误差上找到更合适的裁剪点。注意MSE校准在较新的onnxruntime版本里才可用,版本旧的话优先用Percentile替代。
extra_options里的四个开关是排错时最常动的。ActivationSymmetric设为False是因为非对称量化能让激活的zero_point落在实际范围内,SwinT这种多stage模型收益更明显。MatMulConstBOnly=True表示只量化权重为常量的matmul,不把激活强行量化后经过Dequantize再进matmul,减少INT8和FP32之间的频繁切换开销,这个开关对SwinT尤其重要,否则导出的模型会塞满大量Dequantize节点。QDQOpTypeFallback=True则允许量化器在遇到不支持的算子时自动回退FP32,而不是整个模型量化失败,这也是DeiT/SwinT能一次跑通的关键。
量化完成后,会生成一个deit_base_224_int8.onnx。这个文件里既包含INT8权重,也保留了大量FP32算子,整体文件大小比FP32模型减少30%~50%,同时推理延迟明显下降。到这一步,PTQ量化主体流程就算走完了,接下来要验证精度和性能。
5. 量化后精度崩了怎么排查:三个模型的翻车现场与参数避坑
5.1 精度对比:把INT8结果和FP32结果摆在同一张表
量化做完第一件事,是把FP32 ONNX和INT8 ONNX在同一个验证集上跑一遍top-1精度,三个模型一起对比。我直接写一个统一的评测函数,避免三个模型来回换代码:
import onnxruntime as ort import numpy as np import torch def evaluate_onnx(onnx_path, loader): sess = ort.InferenceSession(onnx_path, providers=['CPUExecutionProvider']) correct = 0 total = 0 for images, labels in loader: logits = sess.run(None, {'input': images.numpy().astype(np.float32)})[0] preds = np.argmax(logits, axis=-1) correct += (preds == labels.numpy()).sum() total += labels.size(0) return correct / total # 验证集至少500张,类别覆盖要均衡 fp32_acc = evaluate_onnx('deit_base_224.onnx', val_loader) int8_acc = evaluate_onnx('deit_base_224_int8.onnx', val_loader) print('FP32 top1: {:.4f} INT8 top1: {:.4f} diff: {:.4f}'.format( fp32_acc, int8_acc, fp32_acc - int8_acc))对比结果通常是:ViT-B/16在ImageNet验证集上,FP32约81.2%,INT8约79.9%,差值1.3个点左右;DeiT-B/16从81.8%到80.5%附近,差值约1.3点;Swin-B从83.5%到81.4%左右,差值约2.1点。这些数值会随timm版本和量化参数有小幅浮动,但整体趋势是SwinT差值最大,主要来自window attention在不同stage间数值范围反复跳动,这不是异常,是这类结构的固有特性。如果你的模型在专用数据集上差值超过了上面这些参考值一倍以上,先别怀疑模型结构,直接检查校准数据和量化配置。
注意不要在PyTorch FP32和ONNX INT8之间直接比,要先确认PyTorch导出成FP32 ONNX的精度损耗在0.1%以内,否则差值里混入了导出的损耗,排查时会把问题推到量化头上。我见过不止一个人为了这0.5%的导出误差折腾了一整天。
5.2 必调的四个量化参数顺序:校准方法、通道粒度、对称模式、回退层
精度不达标时,参数调整有固定顺序,经验是先把误差大头解决,再调细节。
第一步把calibrate_method从MinMax换成MSE。MinMax只看最大值和最小值,对ViT里长尾分布的激活极不友好;MSE会用校准数据找整体重构误差最小的裁剪点,往往一次就能拉回0.5~1个点。第二步确认per_channel=True,这影响的是所有Linear和Conv的量化粒度,ViT的embedding宽度大,per-tensor会让不同通道用同一个scale,精度掉得很厉害。第三步把ActivationSymmetric从True改成False,给激活层一个非对称的zero_point,让量化区间不必覆盖到0,SwinT和DeiT这一步基本能再救回0.2~0.5个点。第四步打开QDQOpTypeFallback,让LayerNorm、Softmax这些不好量化的算子自动回退FP32,这是最后的兜底手段。
如果调整完四步仍然差值大于2个点,该做算子级回退:用nodes_to_exclude把LayerNormalization、或者attention内部特定的MatMul保留为FP32。下面是一个获取LayerNorm节点名的做法:
import onnx model = onnx.load('swin_base_224.onnx') ln_nodes = [n.name for n in model.graph.node if n.op_type == 'LayerNormalization'] exclude = ln_nodes[:4] # 先排除前几个stage的LayerNorm看效果 quantize_static( model_input='swin_base_224.onnx', model_output='swin_base_224_int8.onnx', calibration_data_reader=calib_reader, quant_format=QuantFormat.QDQ, per_channel=True, activation_type=QuantType.QUInt8, weight_type=QuantType.QInt8, calibrate_method=CalibrationMethod.MSE, nodes_to_exclude=exclude, extra_options={'MatMulConstBOnly': True, 'QDQOpTypeFallback': True}, )注意这里排除的是LayerNorm节点,不是它后面的Linear。被排除的节点在INT8模型里仍以FP32方式运行,和QDQOpTypeFallback相比,这个手段可以精确到单个stage、单个节点,排查哪个层在捣乱时很有用。排除掉前四个LayerNorm后,如果精度明显回升,说明是浅层LayerNorm的统计偏移主导了误差,再逐个往外排除,直到找到能接受的精度与加速比的平衡点。
5.3 三个常见坑:从现象到原因的排查手册
第一个坑是ViT量化后准确率掉5%以上,现象非常直接,INT8的top1比FP32低一大截。原因多半是校准数据没做好,比如校准集只有几十张、或者全是同一类别的图,导致激活统计严重偏置。解决方式是先检查校准集数量和类别覆盖度,至少300张且覆盖所有大类别;再把校准方法换成MSE,这两个动作能解决掉大部分“无脑量化”造成的问题。校准集里混入了过多某个固定分辨率或固定背景的图,也会让LayerNorm的scale偏向单一分布。
第二个坑是DeiT导出时报“Could not export Python operator”,现象是torch.onnx.export在forward中途中断。原因是timm的DeiT在forward里走了distill分支,distill head和classifier都是自定义Python算子,onnx导出器不识别。解决方式是导出前用exportable=True,并确认model.head_dist没有被使用;timm的deit模型在distill分支处理上不同版本差异很大,我固定会在导出前打印一次model.head和model.head_dist两个属性,确认head_dist不存在或已被替换。如果还在报错,就用forward hook把distill分支的输出截住,只让classifier输出参与trace。
第三个坑是SwinT INT8推理反而比FP32慢,现象是延迟从30ms涨到45ms,量化加速变成了量化减速。原因出在QDQ格式的Dequantize算子过多,尤其是window attention内部的matmul,每个窗口都要做一次重新量化,算子调度开销高于INT8节省的乘加时间。解决方式是先确认量化模型里Dequantize节点密度,如果密度超过20%,设置extra_options里的MatMulConstBOnly=True,让激活路径不要反复量化;再不行就把attention内部的matmul节点加入nodes_to_exclude,让注意力计算保持FP32,其余路径保持INT8。这样牺牲一小部分加速比,换回正常推理速度。
6. 进阶验证:测真实加速比、判断是否该上混合精度
量化加速值不值得做,最终要看真实延迟而不是模型大小。用onnxruntime自带的方式做一次稳定延迟测试:
import onnxruntime as ort import numpy as np import time def measure_latency(onnx_path, input_array, runs=50): sess = ort.InferenceSession(onnx_path, providers=['CPUExecutionProvider']) for _ in range(5): # 预热,把内存池和线程池拉起来 sess.run(None, {'input': input_array}) times = [] for _ in range(runs): start = time.perf_counter() sess.run(None, {'input': input_array}) times.append(time.perf_counter() - start) return np.median(times) * 1000 # 单位ms inp = np.random.randn(1, 3, 224, 224).astype(np.float32) print('FP32 ms:', measure_latency('deit_base_224.onnx', inp)) print('INT8 ms:', measure_latency('deit_base_224_int8.onnx', inp))用中位数而不是平均值,能去掉系统调度抖动带来的干扰,尤其是跑在共享CPU环境里时,偶发的调度延迟会让均值严重失真。跑完看两个指标:延迟下降比例有没有超过1.5倍,精度损失有没有超过1.5个点。两者都满足,这个PTQ量化就是值得上的。
如果INT8比FP32只快了1.2倍,说明碎算子太多,这时候我习惯的做法是把SwinT的attention内部matmul显式回退FP32,再测一次;ViT和DeiT则优先把LayerNorm回退。如果回退后延迟又掉回去了,就去检查校准数据是否真覆盖了部署场景,再不行才考虑QAT。另一个判断混合精度的依据是打开INT8模型graph,统计Dequantize节点的密度:如果占据总节点数20%以上,说明量化算子切换开销已经压不住收益了,这个模型更适合按算子做混合精度,而不是全量INT8。
我自己第一次部署SwinT时就被“INT8反而更慢”这个现象骗过,当时花了一整天调线程数和batch size,最后打开graph一看才明白是Dequantize节点太多。后来固定先看Dequantize密度再决定是否回退,省了很多不必要的调参时间。量化加速这件事,跑通流程只是起点,摸清自己的模型在哪里量化和在哪里回退,才是真正吃透PTQ的地方。希望帮到你。
本文还有配套的精品资源,点击获取