几年前我在做图像分类项目时,还在一门心思调ResNet的深度和宽度,那时候如果有人告诉我,NLP那套Transformer架构会“反攻”CV,甚至成为主流,我大概率是不信的。但后来的事情大家都知道了:ViT(Vision Transformer)凭借一套纯注意力机制,在ImageNet上杀进了SOTA行列;紧接着Swin Transformer又用一套分层+窗口的设计,把Transformer的通用性和效率一起拉高,连目标检测、分割这类密集预测任务也没放过。这篇博文就好好梳理一下Transformer在CV里的进化脉络,从原理到代码,从不足到突破,尽量把每个关键决策背后的“为什么”讲透,文末附上可以直接改来用的PyTorch代码。
1. 从NLP到CV:Transformer凭什么跨界
1.1 NLP的成功给CV带来了什么启示
2017年《Attention Is All You Need》提出Transformer时,大家的注意力都在机器翻译上。它的核心卖点简单粗暴:不依赖循环结构,直接用自注意力(Self-Attention)建模序列中任意两个位置的关系。这意味着长距离依赖不再是难题,而且整个计算可以高度并行化,训练效率远超当时的LSTM、GRU。
这个设计在NLP领域迅速开花结果,从BERT到GPT系列,预训练大模型成为了标配。而CV这边,很长一段时间仍然是CNN的天下。卷积核天然带有局部性和平移等变性,这让CNN在图像任务上非常高效,也符合图像的某些固有属性。但CNN也有天花板,一个很尴尬的问题是:卷积核的感受野是有限的,虽然可以通过堆叠层数来扩大,但本质上是“由局部逐步组成全局”,对全局关系的建模不够直接。
Transformer的出现给了CV研究者一个全新的视角:既然句子里的词可以通过注意力机制两两交互,图像为什么不行?图像也可以被看成一组像素或一组patch的序列,patch与patch之间的关系,类比词与词之间的关系。于是,用Transformer做图像分类就成了一个自然的想法。
1.2 ViT出现前的那些过渡尝试
在ViT之前,其实已经有不少人尝试把注意力机制引入CV,比较典型的是SENet、Non-local Network。SENet对通道维做注意力,Non-local Network则是对空间位置做全局建模。这些工作有一个共同点,它们不是“取代”CNN,而是给CNN“加装”注意力模块,起到锦上添花的作用。
这种思路在当年很务实,因为CNN的归纳偏置非常强,直接在完整图像上跑自注意力计算量会爆炸。Non-local虽然能建模长距离依赖,但计算复杂度是O(N^2),N为空间位置数量,比如输入是56x56特征图,N=3136,这个复杂度根本不敢往深了堆。
所以ViT真正的贡献不在于“第一个把注意力用在图像上”,而在于它给出了一个极简且有效的方案:把图像切成patch,用标准Transformer Encoder来处理,完全不需要卷积。这个“减法”做得非常彻底,也正因为足够简洁,才让后续大量基于Transformer的CV模型得以快速发展。
2. ViT精讲:把图像切成“单词”
2.1 Patch Embedding是怎么把图像变成序列的
ViT的输入处理非常直接。假设输入图像是224x224x3,设定patch size为16x16,那么图像会被切分成(224/16)^2 = 196个patch。每个patch展平后是一个16x16x3=768维的向量,这个维度刚好可以和Transformer的隐层维度对齐。
但这里有个细节值得注意:直接把展平的像素向量送进Transformer,等于完全抛弃了像素之间的局部空间结构。作者的做法是加一个可训练的线性投影层(Linear Projection),把每个patch的768维向量映射到D维的embedding空间,这个过程和NLP里词嵌入(Word Embedding)本质上是一回事。代码实现简洁到让人意外:
import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.num_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): # x: (B, 3, 224, 224) x = self.proj(x) # (B, embed_dim, 14, 14) x = x.flatten(2) # (B, embed_dim, 196) x = x.transpose(1, 2) # (B, 196, embed_dim) return x看到这个Conv2d了吗?用卷积来实现patch切分是一个很巧妙的工程技巧:卷积核大小等于patch size,步长也等于patch size,这样输出特征图的每个位置就对应一个patch的embedding,一次前向就完成了切块和线性映射两件事。
2.2 Class Token与Position Embedding的细节
标准Transformer的输入是序列,输出也是序列。对于分类任务,ViT借鉴了BERT的做法,在序列开头加入一个特殊的可学习class token,这个token不来自任何patch,它最终对应的输出向量经过分类头后就是整个图像的类别预测。
选择class token而不是简单的全局平均池化,作者在论文里做了实验对比,发现两者的效果非常接近。但class token有一个潜在优势:在预训练阶段,模型可以适应不同数量的输入patch,class token充当了“全局信息汇总者”的角色,这对后面做迁移学习和微调更友好。
Position Embedding是另一个容易被忽略但极其重要的细节。由于Transformer本身是置换不变的,如果不加位置信息,196个patch的顺序就没有任何意义,模型会把图像完全打乱看待。ViT使用的是可学习的1D位置编码,形状是(196+1, 768),加在patch embedding和class token上,一起送入Transformer Encoder。
这里有个实操中的细节:如果预训练时用的是224x224分辨率,微调时换成384x384,patch数量会从196变成576,预训练的位置编码就装不下了。常见做法是对位置编码做二维插值(bilinear interpolation),ViT作者在论文里验证了这个做法有效,但我在实际项目中发现,插值后最好还是让模型在新分辨率上再训练一小段时间,否则位置信息会有一段适应期。
3. Swin Transformer精讲:分层与窗口带来的改变
3.1 ViT的短板和Swin的解题思路
ViT在ImageNet这种千万级数据集上确实能发挥出很强的性能,但训练数据不足时,它的效果往往不如ResNet或EfficientNet。原因也很直接:CNN的归纳偏置(局部性、平移等变)是内置的,而ViT把这些先验全抛掉了,所有模式都必须从数据里学。数据量跟不上,模型就“学不动”。
另外一个让人头疼的问题是多尺度。CNN天然有金字塔结构,浅层特征分辨率高、语义弱,深层特征分辨率低、语义强,这非常契合目标检测和分割这类任务。而ViT在所有层都保持全局建模、固定分辨率(比如14x14),虽然语义信息强,但缺少多尺度特征,直接拿来做检测和分割并不顺手。
Swin Transformer的核心设计就冲着这两个问题去的:一是引入分层结构(Hierarchical),让深层特征分辨率逐层降低,像CNN那样形成特征金字塔;二是引入窗口注意力(Window-based Attention),把自注意力限制在局部窗口内,大幅降低计算复杂度。
3.2 窗口注意力、移位窗口和相对位置编码
窗口注意力的思想非常朴素:图像被均匀划分成若干个不重叠的小窗口(比如7x7的patch区域),每个窗口内部独立做自注意力。这样一来,计算复杂度从ViT的O(N^2)降成了O(N/M),其中N是patch总数,M是窗口内的patch数。以224x224输入、patch size 4为例,Swin第一层的窗口注意力计算量比全局注意力少了一个数量级,这是它能堆出深层模型的重要原因。
但窗口注意力有一个天然缺陷:窗口与窗口之间没有信息交流,某个位置的token只能“看到”自己窗口内的内容,感受野被锁死了。Swin的解法非常妙,它引入了一个移位窗口机制(Shifted Window):当前层用规则窗口,下一层就把窗口整体偏移(比如向右下各偏移M/2个patch),这样上一层在不同窗口里的token,在下一层就有机会进入同一个窗口,间接实现了跨窗口信息流动。
交替使用规则窗口和移位窗口的设计还有一个额外的工程优势:每个窗口内的patch数量保持不变,计算效率高,而且大部分子窗口可以并行处理。
再说相对位置编码。Swin里没有用ViT那种大尺寸绝对位置编码,而是给每对像素位置生成一个相对位置偏置(Relative Position Bias),加到注意力分数上。这样做的好处有两个:一是参数量少,二是有更强的平移不变性。代码实践上,这个步骤通常通过一个可学习的偏置表来实现,PyTorch里可以用nn.Parameter预生成表,然后用索引取出来加在注意力矩阵上面,很多开源实现里都有写成bias表的做法。
3.3 分层架构和特征金字塔的复现
Swin的分层结构跟ResNet的stage概念非常像。整体分为4个阶段(Stage),每个Stage内部的Transformer Block数量不同,常见配置如下:
- Swin-T:每个Stage的Block数为[2, 2, 6, 2],通道数为[96, 192, 384, 768]
- Swin-S:Block数为[2, 2, 18, 2],通道数同Swin-T
- Swin-B:Block数为[2, 2, 18, 2],通道数为[128, 256, 512, 1024]
在每个Stage之间,有一个Patch Merging层,作用类似于CNN里的下采样。以第一个Stage到第二个Stage为例,会把2x2邻域的4个patch在通道维度上拼接起来,得到一个4C维的向量,再通过一个线性层压缩到2C维。这样一来,序列长度减少了4倍,通道数增加了2倍,正好呼应了CNN中“空间分辨率减半、通道数翻倍”的设计范式。
这种结构带来的直接好处是:Swin天然可以输出多尺度的特征图。比如输入224x224,经过4个Stage后可以得到56x56、28x28、14x14、7x7四种分辨率的特征,这正好可以直接喂给FPN、PAN等检测分割头部使用。所以Swin在COCO目标检测任务上,不需要像ViT那样额外设计复杂的特征提取逻辑,直接用分层特征就能取得很好的效果。
4. ViT与Swin,一张表看懂核心差异
很多同学问过我,项目里到底选ViT还是Swin,这个问题没有标准答案,但可以借助对比表来判断:
| 对比维度 | ViT | Swin Transformer |
|---|---|---|
| 核心结构 | 标准Transformer Encoder,全局自注意力 | 分层Transformer,窗口自注意力 |
| 序列长度 | 固定不变(如196个patch) | 逐阶段减半 |
| 计算复杂度 | 随分辨率平方增长 | 随分辨率线性增长 |
| 归纳偏置 | 几乎没有,需大量数据 | 引入局部性和平移不变的先验 |
| 多尺度特征 | 不支持/需要改造 | 天然支持,适合检测分割 |
| 适合场景 | 大数据集分类、预训练、多模态 | 中小数据量分类、检测、分割、密集预测 |
| 典型精度 | ImageNet上需要JFT-300M预训练才能碾压CNN | 纯ImageNet-1K预训练即可达到高精度 |
这里想多说一句,很多人看到ViT在ImageNet-21K上预训练后表现很好,就觉得ViT“就是更强”,但实际上这更像是一个“数据规模和模型容量”匹配的问题。如果你手里只有一两万张图片,直接上ViT很可能会过拟合,换Swin-T或者干脆用卷积模型会更稳妥。
另外需要注意,Swin虽然计算量更小,但因为窗口划分、移位和Mask等操作的实现复杂度高,实际训练时往往会引入一些额外开销。在一些较小的数据集上,Swin的训练速度和显存占用未必比ViT有绝对优势,这个我在后面的实战部分会再细说。
5. 实战:用timm库快速实现并微调ViT与Swin
5.1 环境准备与数据组织
实战部分我直接基于PyTorch和timm来演示,timm这个库把ViT、Swin以及大量Transformer变体都封装成了统一接口,非常适合快速验证想法。我用一个花卉分类的小数据集来演示迁移学习,数据集包含5个类别,每个类别约200张图片,训练集占80%,验证集占20%。
先安装依赖:
pip install torch torchvision timm opencv-python scikit-learn tqdm数据加载部分直接用torchvision的ImageFolder,配合标准的预处理流程。这里有个值得注意的点:ViT和Swin在预处理上基本一致,都要求输入归一化到ImageNet的mean/std,且训练时通常用RandomResizedCrop和RandomHorizontalFlip做增强,验证时用CenterCrop。
import torch from torchvision import transforms, datasets train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) val_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) train_dataset = datasets.ImageFolder('data/flower_5/train', transform=train_transform) val_dataset = datasets.ImageFolder('data/flower_5/val', transform=val_transform) train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4) val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)5.2 一键切换ViT与Swin模型
timm里加载模型非常简单,两行代码的事:
import timm # 加载ViT-B/16,在ImageNet-21k上预训练过 # vit_model = timm.create_model('vit_base_patch16_224_in21k', pretrained=True, num_classes=5) # 加载Swin-T,在ImageNet-1k上预训练 swin_model = timm.create_model('swin_tiny_patch4_window7_224', pretrained=True, num_classes=5)这里有两个细节提醒一下:
第一,timm的模型命名非常讲究,比如swin_tiny_patch4_window7_224这串名字的意思是:tiny配置、patch size为4、窗口大小为7x7、输入分辨率224。建议用之前先跑一下timm.list_models('*vit*')或timm.list_models('*swin*')看看有哪些可用变体,避免名字记错。
第二,如果你用的是在ImageNet-21k上预训练的模型,原始分类头是21841类,迁移到自己的5类数据集时,timm会自动替换分类头,这一点不用操心。但要注意,如果切换的模型使用了不同的输入分辨率,比如vit_base_patch16_384,那么你的预处理也要跟着改成384x384,否则性能会打折扣。
5.3 训练流程与损失函数选择
训练逻辑跟传统CNN训练几乎一样,这也正是Transformer模型在工程落地上最方便的地方:换模型结构不影响训练流程。我习惯用AdamW优化器,配上一个简单的余弦退火学习率调度器,这在Transformer类模型上是公认比较稳定的组合。
from timm.scheduler import CosineLRScheduler from timm.optim import create_optimizer_v2 optimizer = create_optimizer_v2(swin_model, opt='adamw', lr=1e-4, weight_decay=0.05) scheduler = CosineLRScheduler(optimizer, t_initial=30, warmup_t=3, warmup_lr_init=1e-6) criterion = torch.nn.CrossEntropyLoss() for epoch in range(30): swin_model.train() for images, labels in train_loader: images, labels = images.to('cuda'), labels.to('cuda') optimizer.zero_grad() outputs = swin_model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() scheduler.step(epoch) # 验证代码省略练习时可以把模型切换成ViT,用完全相同的训练配置跑一轮,你会直观感受到两者的收敛速度和最终精度差异。
5.4 手写一个精简版ViT核心模块
timm虽然好用,但如果你想深入理解ViT的运行机制,手写一遍核心模块会非常有帮助。下面这个精简版实现,把patch embedding、class token、position embedding和Transformer Encoder都串起来了,去掉了dropout和LayerNorm的一些细节,保留主干逻辑。
import torch import torch.nn as nn class VitForClassification(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12): super().__init__() self.patch_size = patch_size self.num_patches = (img_size // patch_size) ** 2 # patch embedding + 展平 self.patch_embed = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) # class token 和位置编码 self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, embed_dim)) # Transformer Encoder encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim * 4, dropout=0.1, activation='gelu', batch_first=True) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02) def forward(self, x): B = x.size(0) x = self.patch_embed(x).flatten(2).transpose(1, 2) cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat([cls_tokens, x], dim=1) x = x + self.pos_embed x = self.encoder(x) x = self.norm(x[:, 0]) return self.head(x)把这段代码跟timm里的ViT对比,你会发现大体框架是一致的,只是timm实现里做了更多工程细节上的优化,比如用2D位置编码插值、Stochastic Depth、LayerScale等。
6. 训练Transformer模型时踩过的坑与调参经验
6.1 学习率与Batch Size的匹配
Transformer模型对学习率非常敏感,这个跟CNN有比较明显的差别。CNN训练时用SGD加大学习率往往也能跑,但Transformer一开始学习率设大了,经常会看到loss不降反升,甚至直接NaN。我的经验是:ViT和Swin用AdamW时,基础学习率设置在1e-4到2e-4之间比较稳,batch size每翻一倍,学习率也跟着适度上调,但幅度不用跟线性缩放完全一致,1.5倍左右即可。
如果出现loss震荡或NaN,第一件事不是调网络结构,而是把学习率降到原来的十分之一,或者把warmup步数拉长。ViT作者在原文里也比较强调warmup的作用,前几轮让学习率从很小的值慢慢升上去,会有效避免早期不稳定的问题。
6.2 数据增强和正则化策略
Transformer在中小数据集上的过拟合问题比CNN更明显。除了基本的RandomResizedCrop之外,我建议再加一点MixUp或CutMix,这两个增强策略对Transformer的提升比CNN更大。原因也比较直观:Transformer是弱归纳偏置模型,样本多样性的补充能显著缓解过拟合。
另外,Weight Decay的设置在Transformer上也要重新调。CNN里常用的weight_decay=1e-4换成Transformer后可以考虑调到5e-2,因为AdamW对weight decay的处理方式与SGD不同,稍大一点的weight decay能起到更好的正则效果。Dropout率、Stochastic Depth的drop rate也都是可以从默认值往上调的,尤其是在数据量不大的情况下。
6.3 Position Embedding插值与迁移学习
我在一个遥感图像分类项目里遇到过这样一个问题:模型在ImageNet上预训练的时候用的是224x224,但遥感图像往往是512x512甚至更大,直接resize到224会丢失大量细节。于是我把模型输入调成了384x384,这时就需要把预训练的position embedding从196个patch插值到576个patch。
实际操作中,用timm的resize_pos_embed函数可以很轻松完成这个操作,但效果好坏还取决于插值后是否给了模型足够的微调轮数。我的经验是,分辨率切换后,至少需要正常微调轮次的1.5倍到2倍,模型才能把位置信息重新学稳。如果在切换分辨率后只微调很少的轮次就评估,精度往往比直接小分辨率还差,这会让很多人误以为“模型不支持大分辨率”。
6.4 Swin的窗口注意力和显存优化
Swin在显存占用方面其实并不比ViT“天然低多少”。窗口注意力虽然降低了计算量,但PyTorch等深度学习框架在实现自注意力时,为了计算梯度,会保存很多中间变量,显存峰值依然不低。如果你训练Swin时爆显存了,有以下几个实用优化方法:
- 开启
torch.utils.checkpoint(梯度检查点),用计算换显存,速度变慢但显存需求能降一半以上。 - 减小batch size并增加梯度累积步数,效果等同于原batch size。
- 用混合精度训练(AMP),这是所有Transformer模型的标配,不仅省显存,还有大概率加速。
混合精度这块,PyTorch原生支持得很好,加一个torch.cuda.amp.autocast()和GradScaler()就行。我在实际项目里跑Swin-L时,如果不开AMP,一张24G显存的卡连batch size 16都跑不了,开了之后可以跑到32甚至更大。
7. 从ViT到Swin之后的进化方向
聊完ViT和Swin,忍不住想简单展望一下这条技术线路的后续发展,因为对我自己选择模型时有很大参考价值。Swin之后,CV里出现了大量基于Swin的设计做进一步优化的模型,比如ConvNeXt就把Swin的设计思想“反哺”给了CNN,用标准的ResNet架构配上Swin的训练策略和数据增强,最终精度能和Swin打平甚至略高,而推理速度更快。这件事很值得深思:很多时候架构本身的贡献也许没有想象中大,训练策略、数据增强、优化器选择这些细节同样关键。
还有一条线是让Transformer和CNN进一步融合,比如一些混合架构把卷积用在stem和降采样阶段,把Transformer用在高层语义建模阶段,这样既保留了CNN的低层纹理建模优势,又拿到了Transformer的全局建模能力。另一个值得关注的方向是把大核卷积重新推到台前,RepLKNet这类工作证明了大卷积核也能获得接近Swin的性能,但部署上对硬件更友好。
从纯工程角度说,选型时不要盲目追新。如果项目里对推理延迟要求严苛,又希望模型效果好,先不要直接上Swin或ViT,先把ConvNeXt、RepLKNet这些“CNN改进型”测一遍。如果数据集特别大、算力也充足,再考虑大规模Transformer模型加自监督预训练的路子,效果上限往往会更高。
最后分享一个小经验。我在实际调试过程中,经常遇到有人问我:“ViT和Swin到底哪个更好?”我的答案通常是:先搞清楚你的任务类型、数据规模、推理约束这三件事,再谈模型选型。分类任务、数据量大、训练资源充裕,ViT完全值得一试;检测分割、中小数据集、有多尺度需求,Swin大概率不会让你失望。至于中间地带的场景,与其纠结理论优劣,不如把你候选的两三个模型都跑一遍小规模实验,看验证集上的收敛速度和最终精度,这个实测出来的结论永远比任何paper里的对比表更有说服力。这篇梳理从原理讲到了代码细节,一些参数和配置也都是我实际跑过的,希望能帮你少走一点弯路。