简介:这是一份基于 PyTorch 的 Vision Transformer(ViT)实现,面向深度学习研究者与工程师,提供从原始 JAX/Flax 权重转换而来的预训练模型,可直接用于图像分类、特征提取以及下游任务微调。压缩包约 173KB,共 35 个文件,其中以 22 个 Python 源码为主,涵盖模型定义、训练、评估、数据加载与配置管理;另有 README、Markdown 说明、requirements 环境文件、YAML 配置及 Notebook 示例等,目录结构清晰,便于二次开发。资源描述了与原始模型相当的结果,支持 ImageNet2012 等数据集,并附带微调与评估脚本,适合有一定 PyTorch 基础、希望快速复现 ViT 论文或将其应用于自身视觉任务的开发者。已有 7800 余人浏览学习,是入门 Vision Transformer 并获取可用预训练权重的实用参考。 我做视觉这两年,最常用的框架就是PyTorch,而ViT相关的项目里,vision-transformer-pytorch这个库我几乎是反复在用。很多朋友第一次接触Vision Transformer(ViT)时,被论文里一堆概念唬住,总觉得这是个很复杂的模型。但如果你上手跑一遍这个库,会发现ViT的思路其实非常直接:把图像切成小块,当成一串“视觉单词”送给Transformer去处理。而且这个项目坐标很明确——Pytorch + 预训练模型,正好是现在视觉任务落地最常用的一套组合。
这个项目解决的核心问题,就是让大家不用重复造轮子。它提供了完整的ViT模型实现,结构清晰、参数可调,并且可以配合预训练权重直接使用。对我个人来说,它最大的价值在于:既能帮新手理解ViT内部到底发生了什么,也能让有经验的工程师在几行代码内完成模型搭建和迁移学习。无论你是打算做图像分类,还是把ViT当作骨干网络接进检测分割框架,这篇文章都适用。
1. ViT模型为什么值得关注:从CNN到Transformer的范式转移
1.1 图像能否直接当序列处理
在ViT出现之前,视觉模型几乎被卷积神经网络(CNN)统治。CNN的核心假设是局部性和平移不变性:相邻像素关系更密切,同一个卷积核在整张图上滑动。这个假设在ImageNet这类中大规模数据集上非常好用,因为卷积天然带有先验,不需要太多数据就能学会。
但Transformer的想法完全不同。它最初在NLP里证明了一件事:只要数据足够多,你不给模型任何结构先验,让注意力机制自己去找全局关系,效果反而可能更好。于是就有了一个自然的问题:图像能不能也当成一个token序列来处理?ViT回答这个问题的方式很直接——把图像切成固定大小的patch,每个patch线性映射成一个向量,再按顺序拼起来当作一个“句子”来处理。
我最早看到这个思路的时候也挺惊讶,整个ViT居然没有一个卷积层,全靠自注意力在建模。也正是这种极简让它在超大数据集上表现惊人。论文里用JFT-300M这种量级的数据集做完预训练之后,ViT在ImageNet上的精度能超过同等规模的ResNet和EfficientNet。所以它不是换了个结构,而是换了一种视觉建模的范式。
1.2 vision-transformer-pytorch解决的三个实际问题
用这个库一年多,我感受到的突出价值有三点。
第一,实现与论文高度对齐,可读性强。它不是把ViT封装成黑盒,而是把patch embedding、transformer encoder、分类头拆成模块,改哪里都一目了然。我在对比不同层数、不同head数量对精度影响的时候,基本就是改几个参数的事。
第二,原生PyTorch,生态无缝衔接。库本身不依赖timm或者更上层的框架,直接就能和你自己的训练管线集成。我经常需要把ViT输出的特征接给检测头或者分割头,这个库的数据流非常透明,改起来很顺手。
第三,社区认可度高,坑少。这个项目在GitHub上star量很大,用的人多意味着踩坑经验多。比如位置编码的维度问题、patch size选择问题,网上一搜就有很多讨论,遇到bug不至于孤立无援。对做工程和做研究的人来说,这种成熟度很重要。
2. 核心架构拆解:ViT到底在做什么
2.1 Patch Embedding:图像是如何变成Token序列的
ViT最核心的一个操作就是Patch Embedding。以最常见的ViT-B/16为例,输入是224x224的RGB图像,patch_size设为16,那么图像会被划分成(224/16)(224/16)=1414=196个patch。每个patch大小为16x16x3,把它展平成长度为768的向量,再经过一个线性投影层映射到768维的embedding空间。
很多人第一次看会疑惑:为什么不直接展平再用全连接?其实这里的线性投影本质就是一个1x1卷积或者一个reshape加Linear,它的作用是让每个patch的原始像素映射到更适合Transformer处理的语义空间。实际操作中,很多实现直接用nn.Conv2d(in_channels=3, out_channels=dim, kernel_size=patch_size, stride=patch_size)来一步完成切patch和投影,效率更高。这也是为什么你会看到有些代码里Patch Embedding层长得像卷积层,但它做的事情其实就是“切块+线性变换”。
这里我想强调一个点:patch size的选择直接影响序列长度。patch越小,序列越长,计算量越大,但细节保留越多。ViT-B/32用32的patch,序列长度只有49+1个token,速度快很多但精度略降。实际工程里,如果显存有限又不想掉太多精度,可以考虑用大patch,或者保持patch不变减少层数。
2.2 位置编码、CLS Token与Transformer Encoder
patch被映射成token之后,接下来的问题很关键:Transformer本身是顺序无关的,它不知道哪个token在图像的哪个位置。ViT的做法是加一个1D可学习的位置编码向量,直接加到所有token的embedding上。这里没有用NLP里常见的2D位置编码,因为论文实验发现1D可学习编码对效果影响不大,但实现更简单。
ViT还在序列最前面插入了一个特殊的CLS token,它的作用和BERT里的CLS一样,用于汇聚全局信息。在Transformer编码若干层之后,模型拿出CLS token对应的输出向量,接一个分类头完成最终分类。我实际中还见过一些改造方案,比如直接对所有token做全局平均池化再分类,效果有时也不差,但标准ViT用的是CLS token方案。
接下来是Transformer Encoder。以ViT-Base为例,包含12层Encoder,每层由多头自注意力(12个head)、MLP(hidden size从768扩展到3072再降回来)、LayerNorm和残差连接组成。值得注意的是,ViT用的是Pre-LayerNorm结构,也就是每个子层(注意力或MLP)之前先做归一化。这个细节影响稳定性,训练大模型时尤其明显。我自己的经验是,这种设计配合较大的学习率也能保持稳定,微调时不容易崩。
2.3 模型规格怎么选:Base、Large与Huge
ViT官方发布了几个规格:Base(86M参数)、Large(307M参数)、Huge(632M参数)。视觉任务里用最多的就是Base,它和ResNet50规模差不多,但效果更好。如果资源充足、任务复杂,Large往往能带来明显提升;Huge则适合在超大数据集上从头训练,一般做迁移学习的用不起。
我选型时通常会先想清楚数据集规模。数据集只有几千张图,直接用Base甚至Small版本,配合在ImageNet上预训练的权重,效果往往比从零训练要稳得多。所以这里就引出了下一部分的重点:预训练模型到底怎么选、怎么用。
3. 预训练模型的选择与微调实战思路
3.1 三种获取预训练权重的方式对比
标题里特别提到了“带有预训练模型”,这一点其实是很多人最关心的。vision-transformer-pytorch库本身侧重于提供模型结构,而预训练权重通常可以从下面三个渠道获取:
| 获取渠道 | 是否携带官方预训练权重 | 适用场景 |
|---|---|---|
| vit-pytorch库 | 否,只提供模型结构 | 学习结构、自定义改造 |
| timm | 是,ImageNet预训练权重 | 日常分类、工程落地 |
| HuggingFace transformers | 是,Google官方权重 | 研究复现、需要官方预处理 |
我个人的建议是:追求简单就直接用timm,一行代码搞定下载和加载;追求跟原论文对齐就去HuggingFace。用timm加载预训练权重的方法非常直接:
import timm model = timm.create_model("vit_base_patch16_224", pretrained=True, num_classes=1000) model.eval()如果你要做的分类任务类别数不是1000,可以直接在create_model时指定num_classes,timm会自动把最后的分类头替换成对应数量,微调时非常省事。另外timm还提供了丰富的数据增强策略、EMA等训练工具,对训练精度的提升很友好。
如果你更希望复现原论文的预处理流程,可以用HuggingFace的transformers:
from transformers import ViTForImageClassification, ViTImageProcessor model = ViTForImageClassification.from_pretrained("google/vit-base-patch16-224") processor = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224")这里的processor封装了标准化和尺寸调整逻辑,拿过来就能用,不用自己纠结normalize参数。不过要注意,transformers加载这种大权重时会从HuggingFace服务器下载,虽然大部分时候没问题,但偶尔会卡住;这时候可以设置环境变量HF_ENDPOINT=https://hf-mirror.com,用国内镜像加速下载。
注意:预训练权重的下载只是第一步,真正决定模型效果的是后续微调策略。
3.2 微调策略:冻结层、学习率与数据规模
拿到预训练模型后,第一个选择就是:冻结还是不冻结。如果数据集比较小、只有几千张,我建议冻结前11层Encoder,只微调最后一层和分类头。实操中,Encoder底层学习到的都是一些基础纹理、边缘特征,这些对于任何视觉任务都是通用的,不需要重新学。冻结之后可以大幅减少显存占用和训练时间,加快收敛。
如果数据集中等(几万张),或者目标域和ImageNet差异很大(比如医学影像、卫星图),我建议全部微调,但把学习率调低一些。常规的做法是整体学习率设0.0001左右,分类头学习率可以稍微放大到0.001。ViT对学习率比较敏感,一开始就用太大学习率很容易出现loss震荡甚至不收敛。
另外一个小技巧是:如果显存允许,可以先把模型在较大分辨率(如384x384)上微调几轮,效果通常会比224更好。因为ViT没有卷积的局部先验,更大分辨率意味着更多patch,更多细节。代价就是序列长度变长,显存和速度都翻几倍。我自己做细粒度分类时,这个技巧带来的精度提升很明显。
4. 实操记录:从安装到自定义数据集微调
4.1 环境准备和安装
先说我实测过的环境组合:Python 3.10,PyTorch 2.1以上,CUDA 11.8,显卡是RTX 3090。PyTorch的安装直接用官方命令就行,如果下载慢可以把pip源换成清华源或者阿里源。装完之后装依赖:
pip install torch torchvision timm pip install vit-pytorch后面这个vit_pytorch就是lucidrains的版本,也是很多博客里提到的vision-transformer-pytorch的PyTorch实现。不过我想特别提醒一句:vit_pytorch这个包默认不携带预训练权重,它的作用是快速构建模型结构,你想直接跑预训练推理,还是要配合timm或者transformers。很多人一开始没搞清楚这点,装了个vit-pytorch,然后发现模型是随机初始化的,以为库有问题。
4.2 快速推理:用预训练权重分类一张图
下面这段代码是我在项目里验证一张新图时常用的模板,基于timm实现,简单可靠:
import torch import timm from PIL import Image from torchvision import transforms device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = timm.create_model("vit_base_patch16_224", pretrained=True) model = model.to(device) model.eval() transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) img = Image.open("test.jpg").convert("RGB") x = transform(img).unsqueeze(0).to(device) with torch.no_grad(): out = model(x) prob = torch.softmax(out, dim=1) top5 = torch.topk(prob, 5)输出top5之后,去查一下ImageNet的类别索引表就能知道模型预测的是什么类。这里最常踩的坑有两个:一是忘了convert("RGB"),导致灰度图或RGBA图报错;二是忘了加batch维度,unsqueeze(0)少写就报维度错误。我刚开始跑的时候就在这两处浪费过时间。
4.3 自定义数据集微调
假设现在你有一个10类的自定义数据集,目录结构大概是train/class1、train/class2这样。用torchvision的ImageFolder读进来,然后替换最后的分类头,就可以开始微调:
import torch import torch.nn as nn import timm from torchvision import datasets, transforms from torch.utils.data import DataLoader model = timm.create_model("vit_base_patch16_224", pretrained=True, num_classes=10) model.to(device) train_transform = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) train_ds = datasets.ImageFolder("train", transform=train_transform) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.05) criterion = nn.CrossEntropyLoss() for epoch in range(10): model.train() for x, y in train_loader: x, y = x.to(device), y.to(device) loss = criterion(model(x), y) optimizer.zero_grad() loss.backward() optimizer.step() print(f"epoch {epoch}, loss {loss.item():.4f}")这段代码只是最基础的训练循环。实际项目里我还会加warmup、余弦退火、Mixup和数据增强。ViT在中小数据集上容易过拟合,所以增强策略比CNN时代更讲究。还有一个细节:AdamW的weight_decay我习惯设0.05,这是从DeiT论文里来的经验值,实测比0.01稳定。
5. 避坑指南:使用ViT时最容易踩的坑
5.1 输入尺寸和归一化必须匹配
ViT模型对输入尺寸非常敏感,这一点比CNN严格得多。timm里的vit_base_patch16_224要求输入224x224,如果你给它喂512x512的图,patch数量就变了,位置编码的维度对不上,通常运行到模型内部就会直接报维度不匹配的错误。我的建议是,把预处理统一封装成一个函数,Resize到224,归一化参数就用ImageNet默认的mean和std,不要自己随便改。另外,一旦换了数据集,别忘了重新检查normalize参数是否匹配——尤其是医学图像或者红外图像,它们的像素分布和自然图像差别很大。
5.2 预训练权重下载失败怎么办
这个问题在HuggingFace的权重上尤其常见。下载到一半断掉、网络超时、缓存损坏,都会导致无法加载。我遇到这种情况一般分两步排查:先看报错是不是SSL或者超时,如果是,就说明是网络问题,可以设置HF_ENDPOINT=https://hf-mirror.com再重新拉取;如果报错是键名不匹配,那大概率是模型结构定义和权重来源版本不一致,比如用了patch16的定义去加载patch32的权重。权重和模型结构的匹配相当重要。
提示:HuggingFace的缓存目录通常在
~/.cache/huggingface,删掉对应模型的缓存再重新下载,可以解决很多奇怪的加载问题。下载慢但没报错时,多试几次或者手动下载后放到缓存目录也行。
5.3 显存不够用先别急着换显卡
ViT虽然参数不算特别多,但自注意力的计算复杂度是序列长度的平方。224分辨率下197个token还好,一旦输入到448x448,token数变成(448/16)^2+1=785,计算量增长非常明显。如果显存爆了,最直接的思路是减小batch size,或者把patch_size从16改成32。还有一个实用技巧是开启梯度累积,用多个小batch累加梯度模拟大batch,效果能接近但省显存。
训练速度慢的话,优先检查是不是数据加载瓶颈。num_workers调大,或者用pin_memory=True,经常能把GPU利用率拉满。我见过很多新人把num_workers默认0跑,GPU利用率低得可怜,改到4或8之后速度立竿见影。相比一上来就换卡,这招划算得多。
5.4 位置编码与输入分辨率不匹配
如果你想在384x384或更大的分辨率上微调,直接用vit_base_patch16_224的权重会报位置编码维度不匹配。因为预训练权重的position embedding是1x197x768,而384x384对应的是1x577x768。解决办法是插值调整位置编码的尺寸。timm里有些版本支持img_size参数或者在create_model时指定img_size=384,但不是所有实现都自动处理。如果要手动插值,可以这样:
import torch from vit_pytorch import ViT model = ViT(...) pos_embed = model.pos_embedding # 形状 [1, 197, 768] new_pos_embed = torch.nn.functional.interpolate( pos_embed.permute(0, 2, 1).unsqueeze(0), size=(577,), mode="linear", align_corners=False, ).squeeze(0).permute(0, 2, 1) model.pos_embedding = torch.nn.Parameter(new_pos_embed)不过说实话,如果只是做普通分类任务,我不建议手动插值,直接用timm里带384后缀的模型(比如vit_base_patch16_384)会省心很多,位置编码部分timm已经处理好了。
把上面这些坑都踩过一遍之后,我对ViT的理解反而更深了。说实话,ViT的门槛并不在模型本身,而在于各种细节:patch怎么切、位置编码怎么处理、预训练权重怎么融合、微调参数怎么调。把这些细节弄明白,它就是非常趁手的视觉骨干网络。如果你刚开始接触ViT,我的建议是先按着代码把模型结构打印出来,一个个模块核对,再跑通预训练推理,最后再上自己的数据。这个过程走一遍,比干啃论文有用得多。
最后分享一个我自己的习惯:任何新项目要采用ViT,我都会先用Base模型加ImageNet预训练跑一个baseline,确认任务可行之后,再根据显存和精度需求决定要不要换Large或者调分辨率。不要一开始就上大模型,否则调参和排错的成本会高到让你怀疑人生。希望这篇内容对你有帮助,有问题欢迎留言交流。
本文还有配套的精品资源,点击获取