简介:面向深度学习开发者的量化加速实战资源包,专注解决ViT、DeiT、SwinT等Vision Transformer系列模型在推理阶段计算量大、难以部署于资源受限设备的问题。压缩包共15个文件,包含14个Python脚本与1个Markdown说明文档,整体大小仅41KB,结构紧凑。14个Python脚本覆盖了模型定义、量化层实现、整数校准、数据加载、网络封装与多组测试脚本等完整模块,Markdown文档则提供配置说明与量化流程指引,方便开发者快速定位所需内容。目前已有196人学习/下载该资源,适合有一定深度学习基础、熟悉PyTorch并希望掌握PTQ后训练量化技术的算法工程师或研究人员。资源内不仅给出可直接运行的ViT/DeiT/SwinT量化模型,还通过流程教程详细拆解了量化原理、校准方法与加速要点,并借助清晰的项目源码展示从搭建环境到性能评估的完整链路,便于读者结合自身任务二次开发,在边缘端或低算力场景中获得显著推理加速收益。
1. 量化加速 ViT 为什么值得试:几百张图就能把推理压到原来的三分之一
手上已经有一个训练好的 VisionTransformer,想在 GPU 或 CPU 上把延迟压一压,但不想为了加速重新训一个模型——这是做推理优化的人最常遇到的场景。对 VisionTransformer 做 PTQ 量化加速,就是性价比最高的那条路:不需要反向传播,不需要准备训练标签,拿着预训练模型和几百张和部署场景同分布的数据,跑一轮前向统计就能把 FP32 模型转成 INT8 QDQ 模型,常见设备上算子延迟能降到原来的 40% 到 70%,ImageNet 分类的 Top-1 掉点一般在 1% 以内。这套方案同时覆盖 ViT、DeiT、SwinT 三个主流结构,附带模型加载、量化流程和项目源码的工程组织方式,适合正在做服务端推理优化、边缘端部署或者刚入门模型压缩的算法工程师。它解决的核心问题只有一个:如何用最小的改造成本,把 Transformer 的推理成本真正降下来。
2. Transformer 为什么吃 PTQ:哪些算子量化、哪些必须留在浮点
2.1 Transformer 里的算子分成三类:量化收益大的、能但别动的、根本不能碰的
ViT 系列模型的计算量高度集中在 Linear 层。以 ViT-B/16 为例,QKV 投影、注意力输出投影、MLP 两个全连接层占了整个模型超过 90% 的 FLOPs,这些本质上就是矩阵乘法,是 CUDA 上 Tensor Core INT8 内核和 CPU 上 INT8 内核最擅长处理的算子。所以 PTQ 对 Transformer 有效的第一个前提就是:把 Linear 量化掉,就抓住了绝大部分加速收益。
LayerNorm 是第一个要留在浮点的层。它做的是逐通道归一化,每个通道有自己的均值和方差,量化会把这种逐通道的统计信息压成一个全局的 scale,相当于给不同通道的数据强行套了同一个尺子。经验是:LayerNorm 一旦进量化,后面紧接的 Transformer block 输出分布会整体偏移,分类头直接崩掉。Softmax 同理,它输出范围是 0 到 1 的小数,INT8 在 0 到 1 之间只有 255 个刻度,精度损失太明显,而且 Softmax 本身计算量不大,留在浮点不会影响整体延迟。
GELU 这类平滑非线性比较特殊。它不像 ReLU 那样分布天然规整,量化后掉的点通常比 ReLU 多 0.3% 到 0.5%。常见做法是在配置里把 GELU 也排除出量化列表,让它和 LayerNorm、Softmax 一起在浮点分支里计算。有一个简单的经验法则可以记住:凡是对数值范围极度敏感、且本身计算密度不高的层,全部留在浮点;凡是矩阵乘法,全部给量化器。
2.2 对称量化、非对称量化与三种校准算法的选型逻辑
量化参数的核心是 scale(缩放因子)和 zero-point(零点)。对称量化把零点固定为 0,只需要存一个 scale,INT8 范围是 -128 到 127;非对称量化允许 zero-point 非零,能把浮点分布的任意区间映射到 INT8。Transformer 的激活值有个很好的特性:经过 LayerNorm 和 GELU 后分布整体以 0 为中心对称,用对称量化几乎不会浪费表示范围,而且对称量化在 GPU Tensor Core 上的指令支持更好。所以激活和权重都建议默认用对称量化。
校准算法的本质是用一小撮真实数据去估计激活值的分布范围,然后确定 scale。三种常见算法差别很大:
| 校准算法 | 统计方式 | 适用场景 | 翻车风险 |
|---|---|---|---|
| MinMax | 直接取样本中的最大值 | 分布均匀、无长尾 | 异常值会把 scale 拉大 |
| Percentile | 取 P99.9 或 P99.99 分位 | 存在少量离群点 | 截断比例过大 |
| KL 散度 | 用直方图近似分布,选最终失真最小的阈值 | 多峰、长尾、类高斯分布 | 计算量大一些 |
MinMax 实现最简单,但 Transformer 里经常出现个别极端大的激活值,比如 CLS token 经过多层叠加后某个维度突然冲到几十,此时 MinMax 会把 scale 拉得很大,其他 token 的量化分辨率就被压缩了。遇到这种情况,Percentile 或 KL 散度会更稳,其中 KL 散度是最通用的默认选择。
2.3 ViT、DeiT、SwinT 的结构差异如何影响量化方案
这三种模型都叫 Transformer,但对量化器的影响完全不同。ViT 的全局自注意力里,CLS token 要走完整的 self-attention,它的激活值分布和其他 patch token 有明显差异,per-tensor 统计的 scale 容易被 CLS token 带偏。解决思路是校准阶段多给一些样本,让观测器把 CLS token 的极端值看成常态,或者对注意力输出那一层单独配一个 Percentile 观测器。
DeiT 最特殊的是蒸馏机制。它有两个输出头:分类头和蒸馏头,训练时靠 distillation token 从教师模型学知识。量化评估时两个头都要看,不能只盯分类头。如果部署场景只用分类结果,建议直接把蒸馏头相关 Linear 层从量化范围里删掉,减少不必要的精度损耗。
SwinT 采用窗口注意力和 shifted window 机制,早期 stage 的特征图分辨率大、激活范围波动剧烈,后期 stage 通道数多、激活范围收敛。全局共用一个 MinMax scale 会让早期 stage 的极端值拖累后期 stage 的分辨率。这是 SwinT 量化掉点通常比 ViT 高的主要原因。实践中推荐用直方图观测器做 KL 校准,或者给不同 stage 分配不同的观测器,让每个阶段的量化尺度相对独立。
3. 最小复现流程:用 PyTorch 跑通 ViT/DeiT/SwinT 的 INT8 转换
3.1 环境准备与模型加载:从 timm 拉预训练权重
PyTorch 从 1.13 开始的 torch.ao.quantization 提供了一套基于 FX 图模式的量化工具链,支持把模型 trace 成一张计算图,再往图上插入量化节点。这是目前最通用、对 Transformer 支持最好的开源方案。模型统一用 timm 加载,它把三个结构的预训练权重都管理好了。
import torch import timm # 分别加载三种模型,全部切到 eval 模式 vit = timm.create_model("vit_base_patch16_224", pretrained=True).eval() deit = timm.create_model("deit_base_patch16_224", pretrained=True).eval() swin = timm.create_model("swin_base_patch4_window7_224", pretrained=True).eval() # Lemma: 预训练模型必须切 eval,否则 LayerNorm 和 Dropout 的推理统计不对 for model in (vit, deit, swin): model.requires_grad_(False)逻辑说明:requires_grad_(False)是为了确保校准时不会被自动求导图拖慢,同时也避免误开梯度影响统计。FX 量化要求模型在 trace 时行为确定,eval 模式是硬性前提。
参数说明:vit_base_patch16_224表示 ViT-Base、patch 大小 16、输入分辨率 224;swin_base_patch4_window7_224表示 Swin-Base、patch 大小 4、window 大小 7。如果显存紧张可以换成 tiny 版本,量化流程完全一样。
3.2 校准集构造与激活统计循环
校准的本质是让观测器看到真实分布。建议从部署数据里随机抽 1000 到 2000 张图片,覆盖主要类别和典型场景。数据量不是越大越好,超过 5000 张收益很小,还会拖慢校准时间。
from torch.utils.data import DataLoader from torchvision import datasets, transforms transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) calib_dataset = datasets.ImageFolder("data/calib", transform=transform) calib_loader = DataLoader(calib_dataset, batch_size=32, shuffle=False, num_workers=4) # 校准循环:只跑前向,不做任何反向传播 def run_calibration(model, loader, max_batches=64): model.eval() with torch.no_grad(): for i, (images, _) in enumerate(loader): model(images) if i + 1 >= max_batches: break逻辑说明:校准循环和普通推理几乎没有区别,关键是不调用loss.backward(),也不更新任何权重。观测器会在每次前向传播时自动更新激活值的统计直方图。
参数说明:batch_size=32会和部署时的 batch size 保持一致最好,因为激活的分布统计会受 batch 内数据量的影响;max_batches=64对应 64 个 batch、约 2048 张图,这是一个兼顾时间和稳定性的经验值。shuffle=False是为了让校准过程可复现。
3.3 QDQ 图转换:用 QConfigMapping 控制哪些层量化、哪些层跳过
FX 模式下用 QConfigMapping 定义量化策略。核心逻辑是:全局默认量化 Linear 层,LayerNorm 和 GELU 单独排除。
from torch.ao.quantization import QConfig, QConfigMapping, \ PerChannelMinMaxObserver, HistogramObserver from torch.ao.quantization.quantize_fx import prepare_fx, convert_fx def build_qconfig_mapping(model): qconfig = QConfig( activation=HistogramObserver.with_args( dtype=torch.quint8, qscheme=torch.per_tensor_symmetric), weight=PerChannelMinMaxObserver.with_args( dtype=torch.qint8, qscheme=torch.per_channel_symmetric, ch_axis=0), ) mapping = QConfigMapping() mapping.set_global(qconfig) # 将 LayerNorm 和 GELU 排除在量化范围之外 for name, module in model.named_modules(): if isinstance(module, torch.nn.LayerNorm): mapping.set_module_name(name, None) if "gelu" in name.lower() or "act" in name.lower(): mapping.set_module_name(name, None) return mapping example_input = torch.randn(1, 3, 224, 224) vit_prepared = prepare_fx(vit, build_qconfig_mapping(vit), example_inputs=example_input) run_calibration(vit_prepared, calib_loader) vit_int8 = convert_fx(vit_prepared)逻辑说明:prepare_fx会先对模型做符号 trace,生成 GraphModule,再按 QConfigMapping 的规则往图里插入 FakeQuantize 节点。校准阶段,这些节点只统计数值范围,不会真正把数据转成 INT8。convert_fx才是把 FakeQuantize 节点替换成真正的量化/反量化(Q/D)节点,得到 QDQ 模型。
参数说明:ch_axis=0表示按权重矩阵的第一个维度做逐通道量化,对 Linear 层来说就是每个输出通道单独一个 scale,这是保住权重精度的关键;per_tensor_symmetric让激活值全局共用一个对称 scale,这是和硬件指令集最兼容的配置。
3.4 三种模型的跑通差异:DeiT 删蒸馏头,SwinT 换观测器
上面这段代码对 ViT 可以直接跑通,但换到 DeiT 和 SwinT 要注意几个差异点。DeiT 加载后保留分类头和蒸馏头,量化评估时容易出现一个头精度正常、另一个头掉点严重的情况。如果部署只用分类结果,建议先把蒸馏头从模型里摘掉再量化。
# DeiT: 只保留分类头,丢弃蒸馏头 deit.reset_classifier(0, "") # 实际需按 timm API 处理 # 更稳妥的做法是直接重新构建模型 deit = timm.create_model("deit_base_patch16_224", pretrained=True, num_classes=1000)SwinT 的激活分布多峰,建议把激活观测器从 HistogramObserver 换成 MovingAverageMinMaxObserver,并按 stage 粒度设置不同的观测器。窗口注意力的存在让不同 stage 的激活范围差异很大,一个全局 scale 很难同时满足浅层和深层。常见的工程做法是把 SwinT 的 stage 边界找出来,对每个 stage 单独配一个观测器实例,代价是校准时间略长,但精度能回来 0.3% 到 0.8%。
4. 把量化掉点从 2% 压回 0.5%:校准集、校准算法与敏感层分析
4.1 校准集大小和分布:32 张图跑通的是流程,不是精度
很多人第一次跑 PTQ,图省事只拿 32 张图做校准。流程确实能走通,转换也不报错,但验证集上一测掉点经常超过 2%。原因很简单:32 张图覆盖不了真实分布的多样性,观测器把少量样本的局部特征当成了全局规律,估计出的 scale 是有偏的。
校准集大小和精度的关系大致如下:
| 校准集规模 | 适用阶段 | 典型掉点 |
|---|---|---|
| 32 ~ 128 张 | 验证流程是否能跑通 | 1.5% ~ 3% |
| 512 ~ 1024 张 | 常规场景快速迭代 | 0.5% ~ 1% |
| 2000 ~ 5000 张 | 精度敏感、准备上线 | 0.2% ~ 0.5% |
分布比数量更重要。如果部署场景是夜晚街道监控,校准集里就不能全是白天的图片。校准数据务必和真实推理数据的分布一致,否则 scale 估出来就是错的。常见做法是从线上日志里随机采样推理图片,直接存成文件夹做校准集,这比用公开数据集更可靠。
4.2 校准算法怎么选:MinMax 是默认,KL 是兜底
PyTorch 的 FX 量化支持在 QConfig 里直接切换观测器。MinMaxObserver 最朴素,适合激活值分布均匀的模型;HistogramObserver 用直方图近似分布,然后用 KL 散度选出量化失真最小的阈值。实测里 ViT 用 MinMax 通常也能在 1% 掉点以内,但 SwinT 建议直接上 KL,MinMax 在早期 stage 上经常会翻车。
from torch.ao.quantization import MovingAverageMinMaxObserver # 对激活值长尾明显的模型,用 P99.9 截断更稳 percentile_act = MovingAverageMinMaxObserver.with_args( dtype=torch.quint8, qscheme=torch.per_tensor_symmetric, percentile=99.9, )参数说明:percentile=99.9表示忽略激活分布中最大的 0.1% 的值,把剩余范围映射到 INT8。代价是那 0.1% 的较大值会被截断,但对 Transformer 来说是划算的,因为长尾部分的异常值本来就不是有效信息。
4.3 Per-channel 权重量化与 INT32 bias:默认就打开
权重量化有两个档位:per-tensor 是整个权重张量共用一个 scale,per-channel 是每个输出通道各用一个 scale。Transformer 的 Linear 层权重分布逐通道差异明显,per-channel 能带来 0.2% 到 1% 的精度收益,而且现代推理框架对 per-channel 支持都很成熟。bias 在 INT8 推理里通常以 INT32 形式参与累加器,所以不需要对 bias 做量化,转换工具会自动把 FP32 bias 原样附加到 INT32 累加结果上。
一个值得记住的配置模板是:权重用PerChannelMinMaxObserver+per_channel_symmetric,激活用HistogramObserver+per_tensor_symmetric。这个组合在 ViT、DeiT、SwinT 上都是最稳的起点,不需要额外调参。
4.4 敏感层定位:逐层回退找出拖后腿的那几个层
如果整体掉点超过预期,先别急着换更大的校准集,做一次敏感层分析,找到是哪些层的量化误差在放大。做法很简单:把模型量化后,逐层把某个 Linear 层从 INT8 回退到 FP32,其他层保持量化状态,跑一遍验证集,看精度恢复多少。恢复最多的那几层就是敏感层。
def compute_sensitivity(int8_model, fp32_model, val_loader, layer_names): results = {} for layer_name in layer_names: # 将指定层的权重替换为 fp32 原始权重,并移除 Q/D 节点 set_layer_fp32(int8_model, fp32_model, layer_name) acc = evaluate(int8_model, val_loader) results[layer_name] = acc restore_layer_int8(int8_model, layer_name) # 恢复后测下一层 return results逻辑说明:核心思路是控制变量。每次只让一个层走浮点计算,其他层维持 INT8,测出的精度变化就代表这一层量化带来的损失。敏感层定位的意义在于,后续可以用混合精度方案只对这几个层做回退,其他层继续吃 INT8 的加速收益,这样既保住精度又不牺牲太多速度。
5. 避坑:从校准到部署的 5 条踩坑记录
5.1 校准集用了训练集,验证集精度虚高
做过一次线上模型,校准阶段图省事,直接拿了训练集里随机抽的 2000 张图。跑出来验证集精度掉点只有 0.3%,心里觉得稳了,上线后发现线上真实图片掉点接近 2%。原因是训练集和线上分布存在偏差,校准器被训练集里的“常见模式”带偏了。
解决:校准集必须从部署场景的样本分布里采集。没条件采集的话,退而求其次用验证集,但不要和精度评估用的是同一批图,否则属于用测试集调参数,测出来的数字没有参考意义。
5.2 LayerNorm 量化后分类头整体崩掉
一次把整个模型无差别量化,所有层都套上 INT8,结果 ViT 的分类准确率直接从 81% 掉到 63%。单独检查发现是 LayerNorm 的锅:它对每个通道做归一化,量化后每个通道的 scale 被压缩成一个全局值,归一化输出的数值范围完全失真。第一层 LayerNorm 的输出偏差会逐层放大,最后分类头收到的特征分布已经彻底偏移。
解决:把 LayerNorm 排除出量化范围。在 QConfigMapping 里对 LayerNorm 类型的模块调用set_module_name(name, None),让它完全走浮点计算。这是 Transformer 量化的标准操作,不要试图去对它做特殊量化。
5.3 Softmax 量化后注意力分布噪声变大,掉点 0.5% 到 1%
以为 Softmax 计算量小、量化不量化无所谓,就把它留在了量化列表里。结果看图分类的 logits 分布整体变平,置信度普遍偏低,Top-1 掉了将近 1 个点。Softmax 的输出范围是 0 到 1,INT8 表示这个范围只有 255 个刻度,每个刻度约等于 0.004,这对注意力权重的精度来说太粗了。
解决:把 Softmax 留在浮点。它本身不是计算瓶颈,访存占比远大于计算占比,量化它省不了多少时间,反而会引入噪声。凡是对数值精度敏感的非矩阵运算,全部留在浮点分支。
5.4 SwinT 全校准导致早期 stage scale 被带偏
SwinT 量化后整体掉点 1.5%,比 ViT 明显差。排查发现是全局共享一个观测器的问题:SwinT 早期 stage 特征图分辨率高,激活值波动大,偶尔出现较大的离群值;后期 stage 通道数多,激活范围反而收敛。全局 MinMax 观测器被早期 stage 的极端值撑大了 scale,后期 stage 的分辨率就严重不足。
解决:按 stage 粒度拆分观测器。把 SwinT 的 4 个 stage 分别注册不同的观测器实例,让每个 stage 的 scale 独立计算。校准时间会增加一些,但精度通常能回来 0.5% 以上。另一个选择是全局用 KL 散度观测器,它本身对多峰分布更鲁棒。
5.5 量化后推理速度没变快甚至更慢
费劲转完 INT8,部署到 CPU 上一测延迟几乎没变化,GPU 上甚至还慢了。这个现象很常见,原因通常是三个:一是模型太小,INT8 内核的启动开销和 QDQ 节点的计算开销超过了节省的矩阵乘法时间;二是算子没有完成融合,QDQ 节点没有被折叠进相邻的算子,导致数据在 INT8 和 FP32 之间来回转换;三是目标设备不支持某类 INT8 算子,推理框架自动回退到了 FP32,等于量化了个寂寞。
解决:先跑一次算子 profiler,看有没有算子被标记为“fallback to FP32”。如果有,检查该算子的 INT8 内核在当前推理框架里是否受支持。对模型参数量小于 50M 的 Transformer,别指望 INT8 带来质变,先考虑算子融合和内存布局优化。
6. 部署前的验证与优化:让 INT8 真正在目标设备上提速
6.1 用余弦相似度快速找出崩掉的层
只看 Top-1 精度很难判断问题出在哪一层。更有效的做法是逐层对比:把同一批图片分别通过 FP32 模型和 INT8 模型,计算每一层输出的余弦相似度。相似度低于 0.99 的层就是重点关注对象。这个指标比精度更敏感,能在精度还没明显掉的时候提前暴露风险。
def layer_similarity(fp32_model, int8_model, sample_loader): sims = {} hooks_fp32 = register_activation_hooks(fp32_model) hooks_int8 = register_activation_hooks(int8_model) with torch.no_grad(): for images, _ in sample_loader: fp32_model(images) int8_model(images) for name in hooks_fp32: cosine = torch.nn.functional.cosine_similarity( hooks_fp32[name], hooks_int8[name], dim=-1).mean() sims[name] = float(cosine) return sims逻辑说明:激活值余弦相似度衡量的是两个模型在某一层输出的方向一致性。量化误差会沿层累积,越到深层相似度越低,所以一般看最后一层的结果。如果某个中间层相似度骤降到 0.95 以下,那一层就是敏感层。
6.2 ONNX 导出时的 QDQ 融合与算子检查
QDQ 模型导出 ONNX 时,opset 版本要大于等于 13,这样才能把量化/反量化节点表达为标准的 Q/D 算子。导出后要检查是否形成了“Linear + Q + D”这样的三连结构,推理框架只有在看到这种结构时才会触发 INT8 内核路径。如果导出后发现 Q/D 散落在各处,说明算子融合没生效,需要检查是否有自定义算子打断了融合模式。
6.3 混合 INT8/FP16:只把敏感层送回高精度
敏感层分析做完后,通常只有少数几层是真问题。对这些层做混合精度回退,把它们的 Linear 计算换成 FP16 或 FP32,其余层保持 INT8。实际项目里,回退 3 到 5 个敏感层,精度能恢复 80% 的损失,而延迟只增加 10% 到 15%。这种做法的另一个好处是保留了一个后悔药:如果后续验证发现还有问题,只需要继续扩大回退层列表,不需要重新做整模型量化。
部署验证清单: 1. 校准集分布是否和线上一致 2. LayerNorm / Softmax / GELU 是否已排除量化 3. 敏感层列表是否确定,回退策略是否就绪 4. ONNX 导出后 QDQ 结构是否正确融合 5. 目标设备上跑 profiler,确认没有 FP32 回退算子我现在的习惯是拿到一个压缩需求,先跑一轮敏感层分析再做决策,而不是一上来就改配置。量化加速真正费时间的从来不是转换那几步,而是定位哪一层在拖后腿。这个流程从 ViT 到 SwinT 都是通用的,希望帮到你。
本文还有配套的精品资源,点击获取