ARTICLE DETAIL

资讯详情

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

ViT/DeiT/SwinT部署实战:ONNX Runtime PTQ量化与精度调优指南

ViT/DeiT/SwinT部署实战:ONNX Runtime PTQ量化与精度调优指南 简介面向需要部署视觉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, bits8): 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在对称量化里固定为0ViT的权重一般用对称量化zero_point0激活因为分布不一定过零点用非对称量化更合适。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、SoftmaxViT主流技术路线里一个标准的ViT Block包含LayerNorm、MHA和MLP中间穿插GELU。这三个算子恰好是量化的三个难点几乎每次部署都要为它们单独做配置。LayerNorm是对每个token的整条特征向量做归一化输出被拉到一个动态范围很小的区间。初看数值范围稳定但ViT不同层LayerNorm的输出均值差异很大如果全部用同一个scale去量化浅层和深层之间会出现系统性偏差。更关键的是LayerNorm内部有除法、平方根、减法ONNX里的LayerNormalization算子如果在INT8下计算误差会被后面attention放大。我一般会在量化配置里默认让它回退FP32而不是硬量化。GELU在负半轴有一段软饱和区x0时输出不是完全为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 shiftDeiT在ViT基础上加了distillation tokenforward里除了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个点再考虑上QATPTQ阶段能把浮点和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): # exportableTrue 会把timm内部不稳定的自定义op替换成标准onnx算子 model timm.create_model(model_name, pretrainedTrue, exportableTrue).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_version17, # QDQ格式量化需要opset17 ) 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)这里有两个关键参数。exportableTruetimm在导出时会替换掉fx trace不稳定的自定义算子把部分inplace操作和attention的实现改成标准onnx算子这一步能挡掉大部分“Could not export Python operator”的报错。opset_version17ONNXRuntime静态量化在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的目录结构我直接写一个读图片路径的Datasetfrom 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_size8, shuffleTrue, num_workers4)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_inputdeit_base_224.onnx, model_outputdeit_base_224_int8.onnx, calibration_data_readercalib_reader, quant_formatQuantFormat.QDQ, per_channelTrue, activation_typeQuantType.QUInt8, weight_typeQuantType.QInt8, calibrate_methodCalibrationMethod.MSE, extra_options{ ActivationSymmetric: False, WeightSymmetric: True, MatMulConstBOnly: True, QDQOpTypeFallback: True, }, )逐个说明。quant_formatQDQQuantize-Dequantize格式量化后的模型里每个量化算子前后显式插入Dequantize节点这样算子级回退、混合精度和后续的优化都更灵活。per_channelTrue权重按输出通道各一份scaleViT的Linear层维度高这一步不打开精度损失会明显。activation_typeQUInt8激活用非对称无符号8bit因为经过Normalize和LayerNorm之后激活分布不是对称的weight_typeQInt8权重用对称有符号8bit这是CNN时代留下的经验也适用于ViT的Linear和Conv。calibrate_methodCalibrationMethod.MSE比MinMax对长尾分布更鲁棒ViT的激活在attention score附近有明显长尾MinMax会把动态范围拉宽MSE能在整体重构误差上找到更合适的裁剪点。注意MSE校准在较新的onnxruntime版本里才可用版本旧的话优先用Percentile替代。extra_options里的四个开关是排错时最常动的。ActivationSymmetric设为False是因为非对称量化能让激活的zero_point落在实际范围内SwinT这种多stage模型收益更明显。MatMulConstBOnlyTrue表示只量化权重为常量的matmul不把激活强行量化后经过Dequantize再进matmul减少INT8和FP32之间的频繁切换开销这个开关对SwinT尤其重要否则导出的模型会塞满大量Dequantize节点。QDQOpTypeFallbackTrue则允许量化器在遇到不支持的算子时自动回退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_channelTrue这影响的是所有Linear和Conv的量化粒度ViT的embedding宽度大per-tensor会让不同通道用同一个scale精度掉得很厉害。第三步把ActivationSymmetric从True改成False给激活层一个非对称的zero_point让量化区间不必覆盖到0SwinT和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_inputswin_base_224.onnx, model_outputswin_base_224_int8.onnx, calibration_data_readercalib_reader, quant_formatQuantFormat.QDQ, per_channelTrue, activation_typeQuantType.QUInt8, weight_typeQuantType.QInt8, calibrate_methodCalibrationMethod.MSE, nodes_to_excludeexclude, 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导出器不识别。解决方式是导出前用exportableTrue并确认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里的MatMulConstBOnlyTrue让激活路径不要反复量化再不行就把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, runs50): 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的地方。希望帮到你。本文还有配套的精品资源点击获取
返回列表