ViT PyTorch实战:预训练模型加载与微调全攻略

发布时间:2026/9/2 19:42:59
ViT PyTorch实战:预训练模型加载与微调全攻略 简介Vision TransformerViT是近年来将Transformer架构引入图像分类的代表性工作通过将图像划分为固定大小的patch序列并送入编码器以自注意力机制建模全局依赖这份代码提供该模型的PyTorch实现附带从原始JAX/Flax权重转换而来的预训练权重面向希望在PyTorch中直接复现或微调ViT的研究者与开发者。资源覆盖模型定义、训练、评估、数据加载与检查点管理等完整流程支持ImageNet-2012等数据集并提供微调与评估脚本模型代码与训练逻辑分离便于替换主干结构或损失函数可快速在自定义任务上迭代。压缩包共35个文件以22个Python脚本为核心辅以Markdown说明、依赖清单、环境配置及示例Jupyter Notebook总大小仅173KB代码按训练、评估、模型等模块划分结构清晰、开箱即用。目前已有7846人学习下载既可直接加载PyTorch预训练权重也可下载原始JAX/Flax权重并在线转换为PyTorch格式省去自行转换的繁琐示例notebook还展示了从加载模型到推理的完整流程便于快速上手。对于正在复现ViT论文、对比不同实现或开展下游迁移实验的读者这套代码能显著提高实验效率、少走弯路适合具有PyTorch和深度学习基础、希望深入理解视觉Transformer原理并动手实践的研究者与工程师。 Vision TransformerViT近几年在视觉领域的热度不用我多说从2020年Google提出之后它就彻底改变了大家对图像特征提取的认知方式。但是真正上手时你会发现想快速跑通一个带预训练权重的ViT模型在PyTorch生态里还是有不少小坑的。今天分享的这个vision-transformer-pytorch项目正好解决了这个痛点——它是一个纯PyTorch实现的ViT并且直接提供了预训练模型权重开箱即用。无论你是想拿ViT做分类任务的baseline还是想用它替换Backbone做下游任务这篇博文都能给你一条顺畅的上手路径。我会从ViT的核心工作原理讲起再逐个拆解这个项目的代码结构和实现细节然后重点讲解预训练模型的加载、微调和踩坑记录最后整理一份高频问题排查清单。整个过程会穿插我在实际使用中遇到过的问题和解决方式希望能帮你少走弯路。1. ViT模型结构拆解从Patch到Transformer的思路转变1.1 为什么视觉任务也能用Transformer传统CNN靠卷积核滑窗提取特征每一层感受野有限要堆很多层才能看到全局信息。而ViT的思路很直接把一张图片切成固定大小的patch每个patch拉平成向量然后直接丢进Transformer里做全局自注意力建模。这个思路本质上就是NLP中把句子切成token的翻版图像被当作一串视觉单词。这么做最大的好处是模型从第一层开始就有全局感受野不需要像CNN那样靠深度来扩大视野。代价也很明显——在中小规模数据集上ViT越不过CNN的强归纳偏置这道坎这就是为什么预训练模型对ViT特别重要。你拿ImageNet-21K甚至JFT-300M预训练过的权重做初始化再在自己的数据集上微调效果才能发挥出来。1.2 项目整体的技术选型这个vision-transformer-pytorch项目选用PyTorch作为实现框架从工程角度来说是性价比最高的选择。PyTorch的动态图机制让ViT这种带有复杂控制流的模型调试起来非常顺手而且社区生态成熟后续无论是做分布式训练还是用HuggingFace的transformers做迁移都非常便捷。它预训练模型的训练配置也很有代表性主要采用ImageNet-21K预训练加ImageNet-1K微调的两阶段方案跟Google原版ViT论文保持一致。这也是目前绝大多数ViT开源库的通用做法相当于站在了比较靠谱的baseline上如果你需要跑对比实验这个库的成绩是可信的。2. 核心代码逐段拆解ViT是怎么在PyTorch里实现的2.1 Patch Embedding的实现细节ViT的第一个关键操作就是Patch Embedding它负责把图像切块并映射成向量。具体实现其实比想象中简单就是用一个Conv2d配合适当的kernel_size和stride来完成。下面直接看核心代码import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, C, H, W) x self.proj(x) # (B, embed_dim, H/patch, W/patch) x x.flatten(2).transpose(1, 2) # (B, num_patches, embed_dim) return x这里的num_patches计算要留意224x224的图patch_size16得到14x14196个patch。Embedding维度设置为768和Base级别ViT的配置一致。核心逻辑就是用卷积的滑窗特性天然完成了切块和线性映射两步操作比手动切块再过全连接层高效得多。我在实操中特别注意了一点输入图像的尺寸必须能被patch_size整除否则会报维度错误或者隐式裁剪。所以如果你的数据集图片不是正方形建议在预处理阶段做Resize而不是直接往模型里塞。2.2 Positional Encoding和Class Token的细节ViT还有一个非常关键的组件——Positional Encoding。因为Transformer本身没有顺序概念需要靠位置编码告诉模型每个patch在图像中的相对位置。这个项目采用的做法是直接学习一个可训练的位置编码矩阵。self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim))这里的加1是因为还有一个额外的Class Token。这个Class Token是ViT设计中最有意思的地方——它在输入序列的最前面拼接一个可学习的向量经过Transformer编码后这个位置的输出向量就被当作整张图片的特征表示。这种做法借鉴了BERT中的[CLS] token避免了把所有patch的输出都池化或叠加的额外设计。需要注意的是如果微调时输入分辨率变了patch的数量也会变这时候预训练的位置编码就失效了。常用的解法是插值法重新映射位置编码或者直接用新的随机位置编码微调。实际项目中最简单的做法是保持224x224输入不变直接做Resize。2.3 Transformer Encoder中的Attention实现Transformer Encoder的核心就是Multi-Head Self-Attention。ViT在这个项目中对Attention的实现基本遵循了原版Transformer论文但有一个值得注意的细节在Q、K、V的线性变换后做了维度拆分每个Head独立计算注意力再把结果拼接起来。class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasFalse, attn_drop0., proj_drop0.): super().__init__() self.num_heads num_heads head_dim dim // num_heads self.scale head_dim ** -0.5 self.qkv nn.Linear(dim, dim * 3, biasqkv_bias) self.attn_drop nn.Dropout(attn_drop) self.proj nn.Linear(dim, dim) self.proj_drop nn.Dropout(proj_drop) def forward(self, x): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) return x这段代码里有个细节值得展开qkv nn.Linear(dim, dim * 3)是一次性计算Q、K、V的线性变换而不是分别写三个Linear层。这种做法在工程上能减少矩阵乘法的次数提高显存利用率和计算效率。实际测试下来在相同batch size下合并QKV计算比分开写大概能省10%-15%的训练时间。Attention的scale系数是head_dim ** -0.5也就是除以head_dim的平方根。这个设计是为了防止softmax输入过大导致梯度消失属于Transformer训练的标准做法。很多新手在自定义实现时容易漏掉这个细节会导致训练特别不稳定。3. 预训练模型的选择与加载流程3.1 支持哪些预训练权重规格这个项目提供了多种规格的ViT预训练模型核心参数对比如下模型规格Patch SizeEmbedding维度Head数量Transformer层数参数量典型用途ViT-Base16768121286M通用分类/迁移学习ViT-Large1610241624307M高精度任务ViT-Huge1412801632632M大规模数据预训练我个人的经验是显卡显存在24GB以下时优先选择ViT-Base它基本是绝大多数常规任务的性价比之选。ViT-Large需要比较大的batch size和较长的训练时间才能发挥出效果而且微调超参非常敏感除非你的任务本身很复杂、数据量也足够大否则容易杀鸡用牛刀。3.2 加载预训练权重的两种方式第一种是直接使用项目提供的加载接口。这种方式最简单从官方仓库下载.pth权重文件加载后直接做前向推理或微调即可。核心思路如下import torch from vit_pytorch import ViT model ViT( image_size224, patch_size16, num_classes1000, dim768, depth12, heads12, mlp_dim3072, dropout0.1, emb_dropout0.1 ) checkpoint torch.load(vit_base_patch16_224.pth, map_locationcpu) model.load_state_dict(checkpoint[model] if model in checkpoint else checkpoint)第二种是通过HuggingFace的transformers库来加载。这个项目的权重也适配了transformers的接口如果你的下游任务用的是transformers生态直接用AutoModel.from_pretrained即可from transformers import ViTModel, ViTImageProcessor processor ViTImageProcessor.from_pretrained(google/vit-base-patch16-224) model ViTModel.from_pretrained(google/vit-base-patch16-224)这两种方式我推荐按需求取舍。如果只是想快速跑通并验证ViT在你的数据上的效果用方式一就够了它更轻量没有额外的依赖。如果你的项目涉及多个预训练模型的管理和统一调参用transformers接口会更舒服一些。3.3 权重文件下载慢的解决方案下载预训练权重时最让人头疼的就是网络问题。因为权重文件比较大ViT-Base大约330MB如果直接用外网下载速度很慢甚至经常中断。我自己常用的解决办法是先用国内镜像站下载权重文件然后再手动加载到模型里。像hf-mirror.com这类域名可以加速HuggingFace上的模型下载。下载完成后按照上面代码的方式加载即可。还有一个更稳妥的方案在服务器上配置代理环境变量然后直接使用huggingface_hub的snapshot_download函数支持断点续传from huggingface_hub import snapshot_download snapshot_download( repo_idgoogle/vit-base-patch16-224, local_dir./vit_weights, resume_downloadTrue )如果你只是需要预训练权重而不需要整个仓库的配置也可以直接使用hf_hub_download指定filename下载单个.bin或.pth文件。这在排查问题时特别实用不用把仓库几百MB内容全部拉下来。4. 环境搭建与实战从零开始微调ViT4.1 PyTorch环境配置与CUDA版本对应关系ViT的训练和微调对计算资源要求不低GPU几乎是刚需。在配置环境时PyTorch的CUDA版本选择和显卡驱动紧密相关。这里我建议的匹配思路是先查一下显卡驱动支持的CUDA版本命令行里执行nvidia-smi看右上角显示的CUDA Version。这是驱动能支持的最高CUDA版本不代表你要装这个版本但绝对不能超过它。再根据PyTorch官方版本支持情况安装对应的CUDA运行时。比如PyTorch 2.1对CUDA 11.8和12.1都提供了预编译的wheel包你可以根据自己的需求选择。推荐用Anaconda创建独立的虚拟环境这样不会污染系统自带的Python环境出了问题也方便直接删掉重建。conda create -n vit python3.9 conda activate vit pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118很多人在安装时会遇到下载速度慢到怀疑人生的情况。这里我实测最有效的办法是换用国内镜像源加速pip下载。只需要在pip命令后加一个-i https://pypi.tuna.tsinghua.edu.cn/simple速度能提升好几倍。但要注意PyTorch的官方wheel包在镜像源上可能和--index-url参数冲突所以如果发现镜像源找不到对应的CUDA版本包还是回到官方源配合断点续传工具下载。4.2 数据准备与预处理细节ViT对输入数据格式的要求比较固定。我的经验是在数据预处理阶段就统一好尺寸和数据增强策略这比在模型层面反复调参的影响更大。以224x224输入为例推荐的数据预处理流程如下from torchvision import transforms transform_train transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.4, 0.4, 0.4), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ]) transform_val transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ])训练阶段用RandomResizedCrop做随机裁剪相当于给模型提供不同尺度、不同位置的局部信息这是一种强烈的数据增强手段。验证阶段用Resize到256再CenterCrop到224是因为直接Resize到224会拉伸图片比例导致物体变形影响评测准确性。对于Normalize的均值和标准差ImageNet预训练模型常用[0.485, 0.456, 0.406]和[0.229, 0.224, 0.225]。但你如果直接用这个预处理参数在自定义数据集上微调时其实不会出大问题因为归一化只是把像素值映射到标准分布附近。只是如果数据集的颜色分布和ImageNet差异特别大比如医学影像、红外图像最好重新统计数据集的均值和标准差否则模型初始阶段的loss可能偏高。4.3 轻量级微调策略划分特征层与分类头ViT微调和CNN微调最大的不同在于Transformer层之间是强耦合的直接全量微调虽然效果最好但需要很大的显存和较长的训练时间。如果你资源有限我推荐一个折中方案——冻结大部分编码器层只微调最后几层和分类头。for name, param in model.named_parameters(): if blocks.11 in name or head in name: param.requires_grad True else: param.requires_grad False这种做法的原理是Transformer浅层学到的特征偏向通用边缘和纹理深层特征更接近任务相关的高级语义。冻结浅层可以大幅减少计算量同时保留预训练模型强大的特征提取能力。我实测过在一般中小规模分类数据集上只微调最后两层Transformer和分类头效果能做到全量微调的90%左右但训练显存占用直接降了接近一半。如果是更大规模的数据集比如万级以上的数据量还是建议全量微调否则模型容量不够容易欠拟合。4.4 训练参数配置实战经验ViT微调的超参数配置和CNN有比较大的差异。如果你直接用CNN那套经验比如初始学习率0.01、weight decay 0.0001训练很可能直接发散。原因是Transformer内部使用了Layer Normalization和残差结构对学习率更敏感。我常用的配置如下超参数参数值说明优化器AdamW比SGD收敛稳定对Transformer更友好初始学习率3e-5 ~ 1e-4微调时建议从3e-5起步Weight Decay0.05ViT训练论文中的标准配置Batch Size32~128视显存大小而定尽量大一些Epoch数20~50微调不需要太多防止过拟合学习率调度Cosine Annealing比StepLR平滑收敛更好Warmup Steps500~1000先热身再升高学习率防止初期震荡Warmup这一步特别重要。Transformer里的LayerNorm和残差结构在初始化之后并不是完全稳定的如果一开始就用较大学习率容易让预训练权重产生剧烈扰动而且这个扰动往往不可逆后面很难恢复。我见过不少新手第一次微调ViT时loss直接变成nan排查到最后就是没加warmup。5. 实际运行中的高频问题与排查技巧5.1 维度不匹配与输入尺寸问题这个是我遇到最多的报错类型常见错误信息是size mismatch for pos_embed: copying a param with shape torch.Size([1, 197, 1024]) from checkpoint, the shape in current model is torch.Size([1, 197, 768])。出现这个问题的原因是加载的预训练权重是Large规格1024维而当前模型定义的是Base规格768维。解决方案是严格对齐配置参数不要让模型规格和权重规格产生跨级别混搭。如果是因为图像分辨率改变导致的pos_embed维度变化比如从224变成384解决办法是对pos_embed做双线性插值。下面是一段直接可用的代码def interpolate_pos_embed(pos_embed, new_num_patches): # pos_embed: (1, old_num_patches1, dim) num_extra_tokens 1 cls_token pos_embed[:, :num_extra_tokens] pos_tokens pos_embed[:, num_extra_tokens:] old_num_patches pos_tokens.shape[1] new_dim int(new_num_patches ** 0.5) old_dim int(old_num_patches ** 0.5) pos_tokens pos_tokens.reshape(1, old_dim, old_dim, -1).permute(0, 3, 1, 2) pos_tokens torch.nn.functional.interpolate( pos_tokens, size(new_dim, new_dim), modebicubic, align_cornersFalse) pos_tokens pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) return torch.cat((cls_token, pos_tokens), dim1)这个插值操作在更换输入分辨率时几乎是必须的否则会直接报维度错误。我建议在做任何输入尺寸调整之前先跑一次前向检查维度是否正确不要等到训练中途才暴露问题。5.2 显存不足和训练速度慢ViT在训练时的显存占用比同参数量CNN高不少这主要是因为自注意力计算需要存储完整的注意力矩阵。一个14x14的特征图算attention矩阵还好但当特征图分辨率提升或者batch size增大时显存占用会快速膨胀。缓解显存不足的方法有几个我按优先级排列降低batch size这是最直接有效的方案。如果32都跑不动就降到16、8优先保证能跑通。使用梯度累积模拟更大的batch size。每几个step累积一次梯度再统一更新参数。启用混合精度训练PyTorch原生支持torch.cuda.amp在模型精度影响很小的情况下显存占用可以减少接近一半。检查是否关闭了不需要的梯度。如果冻结了某些层记得设置requires_gradFalse这不仅能减少计算量还能降低这部分中间变量的显存占用。训练速度慢的问题通常和DataLoader的数据加载有关。ViT对输入数据量的要求高如果CPU数据加载跟不上GPU训练速度会出现GPU利用率低下、训练空转的现象。我建议设置num_workers4甚至更高同时开启pin_memoryTrue能明显提升数据吞吐。另一个我踩过的坑是在DataLoader里用了太多的随机增强操作比如RandomResizedCrop ColorJitter RandomHorizontalFlip叠加CPU端的计算开销远大于GPU端的推理时间。这时可以考虑把部分数据增强操作放到GPU上进行或者减少一些不必要的增强操作。5.3 微调过拟合问题中小规模数据集上微调ViT过拟合几乎是必然的除非你做了充分的防止过拟合设计。ViT在大模型里其实没有CNN那么强的正则化能力因为自注意力机制很容易记住训练样本中的高频细节。我实验中比较有效的几个方案增加Dropout与DropPathViT在Transformer层中配置了dropout和droppathstochastic depth训练时随机丢弃一部分层能显著缓解过拟合。微调时建议把dropout设置在0.1左右如果过拟合严重可以调到0.2。标签平滑把one-hot标签从硬标签改为软标签比如target变为(1 - epsilon) * one_hot epsilon / num_classes。这能防止模型对训练集过于自信进而提升泛化能力。epsilon一般取0.1。更多数据增强除了基础增强可以加入RandAugment、Mixup、CutMix等更现代的增强策略。ViT在这些强增强下的表现比CNN更稳定这也是Transformer系列模型的普遍特点。Early Stopping监控验证集loss连续多个epoch不下降就停止训练然后用历史最优的checkpoint做推理。5.4 模型评估与特征可视化微调完成后除了看准确率指标建议再做个特征可视化检查。最直接的方法是提取ViT最后一层Attention的权重可视化注意力图看模型关注的是否是图像中真正有判别性的区域。具体做法拿到最后一层Transformer Encoder输出的Attention权重矩阵对每个Head做平均然后缩放回输入图像尺寸用热力图叠加到原图上。如果训练正常注意力图应该集中在目标物体区域而不是零散地分布在背景上。def extract_attention(model, x): # 前向传播并返回最后一层的 attention 权重 model.eval() with torch.no_grad(): attn_weights model.forward_features(x, return_attentionTrue) return attn_weights注意并不是所有ViT实现都暴露了中间层attention需要确认项目代码是否在Transformer Encoder中保留了attention的输出。如果没有保留就需要手动改一下forward逻辑把attention矩阵从上层传递出来。这一步可视化检查在调参时有非常大的意义。我经常发现模型准确率看着不低比如85%但注意力图乱得离谱说明模型学到的特征并不可靠换一批数据很容易崩。注意力可视化能提前暴露这些问题。6. 将ViT作为Backbone迁移到下游任务的要点6.1 替换Backbone的两种路径ViT的一个关键应用场景是作为特征提取器嵌入更大的任务框架中比如目标检测、语义分割、图像检索。这个项目的模型可以灵活对接这些任务常见路径有两条第一条路径是直接使用模型的输出特征。对于分类任务使用Class Token的输出向量对于密集预测任务如分割使用所有Patch Token的输出再配合上采样头还原分辨率。第二条路径是提取中间层的特征。ViT的Transformer Encoder每一层的输出都可以看作是不同抽象程度的特征图。类似CNN中不同stage的特征你可以把blocks.3、blocks.6、blocks.9、blocks.11这几个代表性层的输出拼接起来构成多尺度特征。这在检测和分割任务中是常规操作。6.2 输入分辨率调整时的注意事项把ViT作为Backbone时任务往往需要更大分辨率的输入。典型例子是语义分割通常要求输入尺寸在512x512以上甚至1024。这时就需要对位置编码做插值。要注意的是插值的位置编码会稍微影响模型性能因为原始位置编码的表达空间被改变了。解决方法是在加载预训练权重后先固定其他参数微调位置编码层几百个step让它适应新的输入分辨率。这一步操作简单但效果显著for name, param in model.named_parameters(): if pos_embed in name: param.requires_grad True else: param.requires_grad False # 用你的数据微调 pos_embed步数不用太多300-500步就够实测下来经过这种位置编码热身之后模型在新分辨率下的性能只下降1%-2%远好于直接冻结使用插值后的位置编码。7. 最终踩坑经验总结做ViT相关项目踩坑最多的往往不是模型本身的原理而是工程适配和超参调优。我个人最大的体会是ViT的代码门槛其实不高但它的训练行为和CNN差别很大千万别把CNN那套习惯直接搬过来。几个最值得牢记的操作点第一预训练权重一定要使用ImageNet-21K版本直接在ImageNet-1K权重上微调自定义数据效果差距能达到3到5个百分点。如果项目提供多种预训练配置优先选择大数据集预训练的版本。第二位置编码插值要谨慎使用。能保持224x224输入就尽量保持224x224不要为了适应数据集分辨率随意改动输入尺寸。如果下游任务非得高分辨率试试先在低分辨率下加载预训练模型再渐进式提高分辨率训练这样比直接在高分辨率下初始化稳定得多。第三混合精度训练几乎必开。ViT的显存占用大计算密度也高混合精度在保持精度的同时能明显降低显存压力还能加快训练速度。现在PyTorch原生的torch.amp已经做得很成熟了直接使用即可基本不需要额外处理梯度缩放。ViT真正要大规模应用到自己的业务场景中建议把下面这套组合拳当成标配数据充分增强加预训练权重微调配合位置编码微调和注意力可视化检查。这一步一步走下来模型效果一般不会差到哪里去。后续如果有时间我还会把这个库扩展到更多下游任务场景到时候再继续跟大家同步。本文还有配套的精品资源点击获取