简介:面向深度学习模型部署与推理优化场景,这份资源围绕视觉变换器及其衍生架构ViT、DeiT、SwinT,提供一套完整的PTQ后训练量化加速方案。压缩包内含15个文件,以14个Python脚本为主,覆盖模型定义、量化校准、量化层实现、数据集处理、测试脚本等核心模块,另有1份Markdown流程说明,整体仅41KB,便于对照学习。已有196人学习下载。资源不仅给出经PTQ量化后的ViT、DeiT、SwinT模型,还通过详细教程与可复现源码,帮助开发者理解浮点转整数、校准策略、性能评估等关键环节,掌握将量化技术落地实际项目的完整路径。对希望在边缘设备或资源受限环境中加速视觉模型推理的开发者,具有较高实用价值。
1. 先说结论:Vision Transformer 部署卡在算力上,PTQ 量化是性价比最高的提速手段
把 ViT 这类 Transformer 视觉模型——具体是 ViT、DeiT、SwinT 三个家族——从训练机搬到推理环境,最常遇到的问题是模型能跑,但延迟压不下去。一张 224x224 的图,vit_base 在普通 CPU 上跑到 80ms 很正常,放到端侧设备更夸张。量化加速的思路,是把网络里的权重和激活从 FP32 压成 INT8,计算量直接少一个数量级,推理时还省内存带宽。PTQ(Post-Training Quantization,训练后量化)是三个方案里性价比最高的:不需要重新训练模型,只需要一小批校准图片过一遍前向,得到每层的数值范围,就能完成量化。接下来我会把 ViT、DeiT、SwinT 的量化原理、可复现流程、还有容易翻车的位置一次讲清,适合正在跑模型训练、正准备做部署优化的算法工程师。
2. 量化原理与选型:PTQ 对 Vision Transformer 到底动了什么
2.1 PTQ 与 QAT 的选型:为什么先做训练后量化
模型量化有两条主流路线:PTQ 和 QAT(Quantization-Aware Training,量化感知训练)。PTQ 是在模型训练完成后做一次性转换,不需要动训练流程;QAT 则是在训练过程中就模拟量化误差,让模型自己适应低比特表示,精度更好,但代价是重新训练。对大部分团队来说,第一步永远是 PTQ,原因很直接:它快,而且能告诉你这个模型值不值得继续投入。
| 对比项 | PTQ | QAT |
|---|---|---|
| 数据需求 | 几百张校准图即可 | 需要完整训练集或蒸馏数据 |
| 训练时间 | 分钟级 | 需要重训多个 epoch |
| 典型精度损失 | 1-3 个点(敏感模型可能更多) | 0.5 个点左右 |
| 适合阶段 | 先验证收益,再决定是否继续投入 | 精度要求极严,或已确认收益 |
我见过不少团队直接上 QAT,模型训到一半发现量化收益并不大,白白烧了算力。所以只要不是精度要求特别苛刻的场景,我都会建议先跑 PTQ。如果 PTQ 量化后的精度损失超过 3 个点,再考虑用 QAT 去补。这个顺序能省下大量调参时间和训练资源。
2.2 量化公式与对称/非对称、per-channel/per-tensor 的取舍
量化的核心是把实数映射到整数区间。映射公式就两个:对称量化是q = round(r / scale),反量化是r = q * scale;非对称量化多一个零点,q = round(r / scale) + zero_point,反量化是r = (q - zero_point) * scale。其中 scale 由数值范围决定,zero_point 让整数和实数原点对齐。把这个公式吃透,后面所有参数选择都能推导出来。
对称量化省掉零点,看起来更简洁,但对分布偏斜的激活不公平。ViT 里 GELU 后面的激活值集中在 0 到 6 区间,用对称量化会把负半轴的整数点全浪费掉,分辨率硬生生砍掉一半。所以默认配置是:权重用对称量化,激活用非对称量化。
| 量化对象 | 推荐方式 | 原因 |
|---|---|---|
| Linear / Conv 权重 | 对称 + per-channel | 权重分布近似对称,逐通道能抓住不同输出通道的尺度差异 |
| 激活 | 非对称 + per-tensor | 激活分布偏斜,zero_point 能减少误差 |
| LayerNorm | 保留 FP32 或 per-channel | 统计量不稳定,直接 INT8 容易崩 |
| Softmax / GELU | 保留 FP32 或高精度 | 非线性算子不是计算瓶颈,保留高精度最省事 |
per-channel 和 per-tensor 的取舍也很关键。权重矩阵是[out_features, in_features],不同输出通道的数值范围差异很大,per-channel 给每个通道单独算 scale,量化误差能小一个数量级。激活则不然,推理时激活 tensor 是动态生成的,per-channel 意味着每个 batch 都要实时统计每个通道的 min/max,开销大且不稳定,所以激活几乎都用 per-tensor。
动量观察器和直方图观察器在这里发挥作用。MinMax 观察器直接用校准期间看到的最大最小值,简单但容易被离群点带偏;Percentile 观察器忽略少数极端值;Histogram 观察器用直方图推断更合理的范围。我一般会先用 MinMax 跑一版,如果掉点超过预期,再换 Histogram 调 bins 数重跑校准,大部分模型都能在这个区间救回来。
2.3 LayerNorm、Softmax、GELU:三个量化敏感点的机理
Vision Transformer 和 CNN 的算子结构差异很大。CNN 里大量是 Conv 加 ReLU,ReLU 输出非负,量化范围很好定;而 ViT、DeiT、SwinT 主要由 Linear 构成,中间夹着 LayerNorm、Softmax、GELU,每个都是量化坑。先看 LayerNorm:它对每个 token 的特征维做标准化,减均值除以标准差。问题出在标准差上——标准差小的时候,除出来的数值会很大,INT8 的[-128, 127]区间根本不够用;而且 LayerNorm 后的 tensor 分布在不同 batch 间变化很大,校准得到的 scale 容易在边界上翻车。这也就是为什么 ViT 系列的量化通常比同规模的 CNN 更难做。
Softmax 的输出是 0 到 1,看起来温和,但中间步骤是 exp 求和再除。INT8 的 256 个刻度去表达 0 到 1 的数值,精度大约只有百分之一。更麻烦的是 exp 之前还要做减 max 的操作,涉及的范围很宽,量化后这一步的误差会直接传导到后续 Linear 层。部署时我一般保留 Softmax 在 FP32,或者至少把分母用 FP32 算,只让最终 Linear 走 INT8,这样处理后的精度损失基本可以忽略。
GELU 在负半轴有一段平滑非线性,不像 ReLU 那样硬截断。量化后,负半轴的小数值容易被压到 0 附近同一个刻度上,输出分布出现系统性偏移。CNN 里 ReLU 量化后基本没感觉,GELU 会明显一点,特别是在深层 block 里误差逐步累积。遇到这种情况,优先把 GELU 这一段保留 FP32,代价只是少部分算子不走量化加速,但对整体精度有稳定效果。
3. 落地路径:套模板把 ViT/DeiT/SwinT 用 PTQ 跑通
3.1 环境准备与模型加载
开始动手之前,先把环境确认好。PyTorch 1.13 之后的torch.ao.quantization提供了面向 FX 图模式的量化接口,推荐直接用 PyTorch 2.x。另外确认目标 CPU 支持 AVX2 指令集——FBGEMM 后端依赖它,不支持的情况下量化后的算子会绕回慢速路径,速度收益基本为零。
import torch import timm torch.manual_seed(0) model_names = { "vit": "vit_base_patch16_224", "deit": "deit_base_distilled_patch16_224", "swin": "swin_base_patch4_window7_224", } model = timm.create_model(model_names["vit"], pretrained=True).eval() print(model)这里有三件事要注意。第一,timm.create_model加载的是预训练权重,可以直接拿来量化,不需要重新训练。第二,一定要先.eval()再量化,因为 Dropout 和 LayerNorm 在 train 和 eval 模式下行为不同,量化统计必须基于推理模式的确定性行为。第三,模型名要和你手上的权重来源对齐,同一个结构在不同训练配置下权重分布会有差异,后面量化精度的表现也会不一样。
3.2 校准数据准备与观察器选择
校准数据是 PTQ 质量的关键。常见做法是从训练集里抽一小部分图片,覆盖所有类别,数量控制在 200 到 1000 张,batch size 设成 16 或 32 跑一轮前向。注意这里不要用验证集做校准,验证集数据在校准时泄漏到量化范围里,最后测出来的精度会虚高。这是 PTQ 里面最容易忽视的翻车点。
import timm.data from torchvision import datasets from torch.utils.data import DataLoader transform = timm.data.create_transform((224, 224), is_training=False) calib_dataset = datasets.ImageFolder("/data/calib_images", transform=transform) calib_loader = DataLoader(calib_dataset, batch_size=16, shuffle=True, num_workers=4) calib_iter = iter(calib_loader) sample_batch = next(calib_iter) example_inputs = sample_batch[0]用timm.data.create_transform而不是自己拼 transform 的原因,是它会自动匹配模型默认的 mean/std 和插值方式。ViT 和 SwinT 的数据预处理细节有差异,手动拼 transform 很容易在数据分布上翻车。example_inputs会在后面的prepare_fx里作为样例输入,用来追踪图结构,所以它的 shape 必须和真实推理输入完全一致,不能随意改。
3.3 prepare_fx 与 convert_fx:PTQ 的最小可跑代码
PyTorch FX 量化是目前最值得抄的模板,核心就三个步骤:准备、校准、转换。prepare_fx在模型图里插入 Observer 节点,前向跑校准数据时统计每层激活的范围;convert_fx把 FakeQuantize 节点替换成真正的 INT8 算子。整个流程如下。
from torch.ao.quantization.quantize_fx import prepare_fx, convert_fx from torch.ao.quantization import get_default_qconfig qconfig = get_default_qconfig("fbgemm") # x86 CPU 用 fbgemm qconfig_dict = {"": qconfig} prepared_model = prepare_fx(model, qconfig_dict, example_inputs) with torch.no_grad(): for images, _ in calib_loader: prepared_model(images) quantized_model = convert_fx(prepared_model) torch.save(quantized_model.state_dict(), "vit_ptq.pt")校准过程中with torch.no_grad()是必须的,量化的统计只依赖前向数值,不需要梯度。qconfig_dict里空字符串""表示全局配置,后面可以按模块类型或模块名覆盖。一个实用的细节是,校准的 batch 数不用太多,50 到 100 个 batch 足够让每层的 min/max 稳定下来;跑太多反而会让个别离群样本占据量化范围,精度更差。
3.4 量化模型的精度与延迟验证
量化完成之后,先验证精度,再测延迟,顺序不能反。精度不达标时测延迟没有意义。
def evaluate_on_val(model, loader, device="cpu"): model.eval() correct = total = 0 with torch.no_grad(): for images, labels in loader: images, labels = images.to(device), labels.to(device) out = model(images) pred = out.argmax(dim=-1) correct += (pred == labels).sum().item() total += labels.size(0) return correct / total from torchvision import datasets val_dataset = datasets.ImageFolder("/data/val_images", transform=transform) val_loader = DataLoader(val_dataset, batch_size=32, num_workers=4) fp32_acc = evaluate_on_val(model, val_loader) int8_acc = evaluate_on_val(quantized_model, val_loader) print("fp32_acc=%.4f int8_acc=%.4f" % (fp32_acc, int8_acc))延迟测试要避免只跑一次。第一次调用有线程池启动、指令缓存预热等开销,建议先跑 10 次预热,再正式测 50 轮取平均值。CPU 推理用 batch size 1 测单张延迟更接近真实业务场景,如果业务是批量处理的也可以测吞吐。
import time def latency_ms(model, sample, repeats=50): sample = sample[:8] # 固定 batch=8 with torch.no_grad(): for _ in range(10): model(sample) start = time.time() for _ in range(repeats): model(sample) return (time.time() - start) / repeats * 1000 print("fp32: %.2f ms" % latency_ms(model, example_inputs)) print("int8: %.2f ms" % latency_ms(quantized_model, example_inputs))这里你会经常看到一种情况:int8 在 PyTorch 里测出来的延迟没有明显下降,甚至更慢。先别下结论说量化没用,更常见的原因是算子没有真正走到 INT8 内核,或者后端选错了。这个坑我在第 5 章展开说。
4. 三个模型的量化差异:ViT、DeiT、SwinT 不能一套参数走天下
4.1 ViT:基准模型,掉点能控制在 1 到 2 个点
ViT 是最干净的 Transformer 结构,Patch Embedding 加若干层标准 Encoder Block,最后接一个分类头。因为没有窗口操作和蒸馏分支,用默认 PTQ 配置跑出来的效果一般都不差。以 ImageNet 验证集为例,vit_base 量化后掉点通常在 1 到 2 个点范围内,超过这个数基本可以断定是校准数据或观察器参数的问题。
如果默认配置掉点超过预期,我会先把观察器从 MinMax 换成 Histogram。具体做法是用 QConfig 手动指定:
from torch.ao.quantization import QConfig from torch.ao.quantization.observer import HistogramObserver, MinMaxObserver qconfig = QConfig( activation=HistogramObserver.with_args( dtype=torch.quint8, qscheme=torch.per_tensor_affine, bins=2048, ), weight=MinMaxObserver.with_args( dtype=torch.qint8, qscheme=torch.per_channel_symmetric, ), )HistogramObserver 的优势在于它不直接用最大值,而是通过直方图选择合适的截断点,能容忍少量离群值。bins 设 2048 是性能与精度之间的一个常规取值,太小会丢失分布细节,太大会让校准变慢但精度提升不明显。ViT 的全局注意力让每层激活的分布比较规律,观察器调整后效果很直接。
4.2 DeiT:蒸馏 token 的量化陷阱
DeiT 的量化难点不在算子,而在模型结构本身。DeiT 在训练时引入了一个蒸馏 token,让模型同时学习分类和蒸馏两个目标。量化时如果不知道这个结构,会在验证阶段选错输出头,看起来精度崩了,实际是拿蒸馏分支的结果去和标签做对比。
在 PyTorch 里,deit_base_distilled_patch16_224这类蒸馏版本模型,forward 返回值可能包含多个张量。量化前先确认输出结构:
out = model(example_inputs) if isinstance(out, (tuple, list)): cls_logits, dist_logits = out print("cls head shape:", cls_logits.shape) print("dist head shape:", dist_logits.shape)部署或验证时,默认取分类头那一支。如果业务需要同时使用两个头,两个分支都要参与量化校准,不能只跑其中一个分支就让模型转换。蒸馏 token 分支的统计特性不同,量化配置要单独确认,蒸馏分支量化后的精度往往比主分类头更敏感。
4.3 SwinT:窗口注意力与 patch merging 的量化误差
SwinT 是这三个模型里量化最需要小心的。它把特征图切成 7x7 的窗口做局部注意力,配合滑动窗口跨窗口通信。窗口边界的 padding 在量化时会引入不连续的数值分布,而这些分布被同一个激活 scale 覆盖,导致量化误差偏大。
patch merging 是另一个隐藏问题。它把相邻 token 拼接后过 Linear 层降维,拼接操作让数值范围直接变大,再经过量化,误差会被放大。所以 SwinT 量化后掉点通常比 ViT 明显,常见在 2 到 4 个点。遇到这种情况,我一般优先对 patch merging 相关的层单独设高精度。
实际经验是,SwinT 用同一套默认配置跑出来精度差,不一定是操作错误,更多是结构特性。可以考虑混合精度——把最后的几个 block 保留 FP32,前面的层走 INT8,可以明显缓解误差累积。但代价是推理速度的提升会打折。量化之前先用 TensorBoard 或脚本看每层的激活范围抖动情况,能更早判断哪些层需要在量化配置里特殊处理。
5. 避坑专题:PTQ 量化 Vision Transformer 的 5 个常见问题
5.1 现象:精度掉了 5 个点以上,校准数据本身有问题
原因:最常见的是用了验证集做校准,或者校准图片的来源和真实部署场景不一致。验证集在校准时信息泄漏,scale 和 zero_point 被调得恰好适配验证集,换到真实数据就崩。另一个常见问题是校准集数量太少,只有几十张,统计出来的 min/max 不稳定。
解决:从训练集里重新抽样,按类别比例均匀抽取,数量控制在 200 到 1000 张。如果模型是跑在某个特定领域的,校准集必须包含该领域的代表样本,不能用通用数据集强行替代。
5.2 现象:量化后输出出现 NaN 或精度集中在最后几个 block 崩掉
原因:Softmax 被量化后,exp 求和这一步的 INT8 精度不够,最后几个 Linear 层拿到的输入已经带了不小误差,误差累积后表现成输出异常。
解决:把 Softmax 模块排除出量化范围。在qconfig_dict里按模块类型单独设置:
qconfig_dict = { "": qconfig, "module_type": [(torch.nn.Softmax, None)], }None表示该模块不量化,保持 FP32。这个改法对 ViT、DeiT、SwinT 都有效,因为它们的 encoder 里都有 Softmax,属于通用修复手段。
5.3 现象:量化后精度温和下降但又不至于崩,LayerNorm 的小数位被抹掉了
原因:LayerNorm 输出范围不稳定,用 per-tensor scale 覆盖全部通道时,一部分通道的小数值差异被压到同一个整数刻度上,深层网络里误差逐层累积。
解决:把 LayerNorm 也排除出量化范围,或者单独给它配高精度的观察器。注意这个操作会牺牲一点推理速度,LayerNorm 本身不是计算瓶颈,但会破坏算子融合的连贯性。如果目标是极致速度,可以先试只保留最后的 LayerNorm 为 FP32,前面的量化掉层继续跑 INT8。
5.4 现象:量化后精度没问题,但推理延迟没有降下来
原因:最常见的有三个。第一,后端选错,x86 CPU 上用 fbgemm,ARM 端侧要用 qnnpack,配置不对算子会走慢速回退路径;第二,部分算子在目标后端不支持 INT8,自动 fallback 回 FP32;第三,模型本身太小,量化算子调度的额外开销抵消了 INT8 的计算收益。
解决:先用 PyTorch Profiler 看算子级别的耗时分布,确认 Linear 层是否真的走了 INT8 GEMM 内核。再换到目标运行时验证——把量化模型导出 ONNX,用 ONNX Runtime 或 OpenVINO 重新跑延迟,这两个运行时对 INT8 算子的融合比 PyTorch 原生路径激进得多。
5.5 现象:SwinT 量化后精度勉强能接受,但速度纹丝不动
原因:SwinT 的窗口切分、移位、padding 这些操作在很多推理后端上没有对应 INT8 算子,运行时只能把整段子图回退到 FP32。量化最核心的 Linear 层反而没有真正提速。
解决:用 ONNX Runtime 的 QDQ 模式执行,观察 graph 结构里有没有大量回退节点。如果回退太多,就把策略调整为“只量化 Linear,其他算子保留 FP32”。还有一个偏经验的做法是把窗口大小从 7 改成 8 的倍数再量化,虽然不保证所有后端都买账,但在部分推理库上能让算子走更规整的高效路径。这个改法会影响模型精度,需要重新校准和验证,不能无脑套用。
6. 进阶技巧:部署提速前,先写好验收脚本和算子级观察
6.1 用验收脚本锁定量化收益
量化项目最常见的失败方式,不是精度崩了,而是没把“收益”这件事讲清楚。精度对比、延迟对比、模型体积对比,应该在一份脚本里一次性跑出来,数据落在同一张表里,作为是否继续投入 QAT 的决策依据。
import time import torch def evaluate(model, loader, device="cpu"): correct = total = 0 with torch.no_grad(): for images, labels in loader: out = model(images.to(device)) correct += (out.argmax(-1).cpu() == labels).sum().item() total += labels.size(0) return correct / total def measure_latency(model, sample, repeat=100, device="cpu"): with torch.no_grad(): sample = sample.to(device) for _ in range(10): model(sample) start = time.time() for _ in range(repeat): model(sample) return (time.time() - start) / repeat * 1000 report = [] for name, m, loader in [("fp32", model, val_loader), ("int8", q_model, val_loader)]: acc = evaluate(m, val_loader) lat = measure_latency(m, example_inputs[:8]) report.append((name, acc, lat)) print(report)验收标准我一般这样定:掉点在 1 个点以内,量化直接上线;掉点 1 到 3 个点,结合业务场景判断是否接受;掉点超过 3 个点,再投入做 QAT 或混合精度调优。没有这份脚本,后面做的一切优化都是凭感觉,出了问题也很难定位是量化配置、后端选择还是数据分布的问题。
6.2 导出 ONNX 用运行时重新验证
PyTorch 原生跑 INT8 的加速有限,真实落地时通常会导出到目标推理后端。导出本身很简单,但验证方法要跟着变:
torch.onnx.export( quantized_model, example_inputs[:8], "vit_ptq.onnx", opset_version=17, do_constant_folding=True, )导出的 ONNX 里带 QDQ 节点,用 ONNX Runtime 加载时,graph 优化阶段会把相邻的 Quantize、DeQuantize 和算子做融合,实际推理速度通常比 PyTorch 原生路径好。注意opset_version至少用 17,低版本对 QDQ 的支持不完整,部分量化信息会在导出时丢失。
按照经验,同一个量化模型在 PyTorch 原生测延迟 60ms,导出到 ONNX Runtime 可能只要 40ms,再到 OpenVINO 或 TensorRT 上又是另一档表现。评价量化收益必须绑定目标后端,也就是最终部署环境里的推理运行时,不能只看 PyTorch 里的数字。
另外,量化后模型还可以继续做通道剪枝或算子融合,但顺序上先量化再剪枝,避免剪枝后重新跑量化流程。我个人的习惯是量化前就写好验收脚本,把精度基线和延迟基线先存下来,量化后的一切改动都对着这份基线对比,合格就继续,不合格就回滚到上一版配置继续调。有一段时间我在测试环境里看到 int8 延迟没变化就以为量化没用,后来换到目标运行时一测,速度提升接近三倍。栽过这个跟头之后才明白,量化加速的判断要放在真实的部署环境里做,不能拿测试环境的结果直接给方案定性。希望这些从原理到踩坑的过程对你有帮助,也希望你在量化 ViT 系列模型时少走这些弯路。
本文还有配套的精品资源,点击获取