ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

ViT/DeiT/SwinT PTQ量化实战:将Transformer推理提速至三倍

ViT/DeiT/SwinT PTQ量化实战:将Transformer推理提速至三倍 简介面向深度学习开发者的量化加速实战资源包专注解决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只需要存一个 scaleINT8 范围是 -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, pretrainedTrue).eval() deit timm.create_model(deit_base_patch16_224, pretrainedTrue).eval() swin timm.create_model(swin_base_patch4_window7_224, pretrainedTrue).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、输入分辨率 224swin_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, transformtransform) calib_loader DataLoader(calib_dataset, batch_size32, shuffleFalse, num_workers4) # 校准循环只跑前向不做任何反向传播 def run_calibration(model, loader, max_batches64): model.eval() with torch.no_grad(): for i, (images, _) in enumerate(loader): model(images) if i 1 max_batches: break逻辑说明校准循环和普通推理几乎没有区别关键是不调用loss.backward()也不更新任何权重。观测器会在每次前向传播时自动更新激活值的统计直方图。参数说明batch_size32会和部署时的 batch size 保持一致最好因为激活的分布统计会受 batch 内数据量的影响max_batches64对应 64 个 batch、约 2048 张图这是一个兼顾时间和稳定性的经验值。shuffleFalse是为了让校准过程可复现。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( activationHistogramObserver.with_args( dtypetorch.quint8, qschemetorch.per_tensor_symmetric), weightPerChannelMinMaxObserver.with_args( dtypetorch.qint8, qschemetorch.per_channel_symmetric, ch_axis0), ) 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_inputsexample_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_axis0表示按权重矩阵的第一个维度做逐通道量化对 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, pretrainedTrue, num_classes1000)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 建议直接上 KLMinMax 在早期 stage 上经常会翻车。from torch.ao.quantization import MovingAverageMinMaxObserver # 对激活值长尾明显的模型用 P99.9 截断更稳 percentile_act MovingAverageMinMaxObserver.with_args( dtypetorch.quint8, qschemetorch.per_tensor_symmetric, percentile99.9, )参数说明percentile99.9表示忽略激活分布中最大的 0.1% 的值把剩余范围映射到 INT8。代价是那 0.1% 的较大值会被截断但对 Transformer 来说是划算的因为长尾部分的异常值本来就不是有效信息。4.3 Per-channel 权重量化与 INT32 bias默认就打开权重量化有两个档位per-tensor 是整个权重张量共用一个 scaleper-channel 是每个输出通道各用一个 scale。Transformer 的 Linear 层权重分布逐通道差异明显per-channel 能带来 0.2% 到 1% 的精度收益而且现代推理框架对 per-channel 支持都很成熟。bias 在 INT8 推理里通常以 INT32 形式参与累加器所以不需要对 bias 做量化转换工具会自动把 FP32 bias 原样附加到 INT32 累加结果上。一个值得记住的配置模板是权重用PerChannelMinMaxObserverper_channel_symmetric激活用HistogramObserverper_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 到 1INT8 表示这个范围只有 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 都是通用的希望帮到你。本文还有配套的精品资源点击获取
返回列表