1. 整体思路拆解为什么Transformer能“跨界”到图像领域先说个结论Vision TransformerViT不是把Transformer原封不动搬到图像上而是把图像“翻译”成Transformer能理解的语言——也就是序列。以前我们做图像分类首选CNN从AlexNet到ResNet靠卷积核一层层提取局部特征再靠堆叠网络深度扩大感受野。这种方式有效但有个隐含假设图像特征是局部的、分层的越靠近输入越关注边缘纹理越靠近输出越关注语义。这个假设在自然图像上基本成立所以CNN一直很强。但Transformer的思路完全不一样。它在NLP里处理的是文本序列每个词是一个token靠Self-Attention机制直接建模任意两个token之间的关系。它没有“局部优先”的归纳偏置而是让模型自己学会该关注哪里。有人就想了如果把图像也切成一个个小方块当作token序列喂给Transformer让模型自己学patch之间谁和谁相关是不是也行ViT在2020年由Google提出实验证明当训练数据足够大比如JFT-300MViT在ImageNet上的表现可以超过同量级的CNN。即便数据量不够通过合适的训练策略和迁移学习ViT照样能打。这篇文章的定位很清楚给刚接触Transformer、想搞懂ViT原理并能自己写出代码的读者。我会从为什么用Transformer做视觉、每个模块到底在算什么一步一步落到PyTorch代码。代码不是调用现成库而是从零搭建一个完整的ViT保证你跑完一遍之后以后再看到各种ViT变体DeiT、Swin、hgformer这类都能看懂门道。2. ViT核心原理把一个图像拆成一句话2.1 图像分块从二维像素到一维Token序列一张图在计算机眼里是一个三维数组形状是[H, W, C]H是高度W是宽度C是通道数。比如224x224的RGB图就是[224, 224, 3]。ViT第一步把图切成一个个小方块论文里叫Patch。比如把每16x16像素划成一个Patch那么在224x224的图上一共可以切出(224/16)² 196个Patch。每个Patch原本的形状是[16, 16, 3]。为了把它当作Transformer的输入token要把它“拉平”成一个一维向量也就是把这16x16x3768个数值全部排开得到一个长度为768的向量。这个操作在深度学习里叫Flatten。这样一来一张图就变成了一个长度为196的序列序列里的每个元素是一个768维的向量。这个结构就跟NLP里的句子一模一样了196个token每个token的维度是768。但还不能直接把原始像素塞进Transformer因为原始像素的数值范围不稳定而且直接在高维原始特征上做Attention效果不好。所以ViT加了一个线性映射层Linear Projection把这个768维的原始向量映射到模型隐层维度DD在论文里通常是768、1024这类数值。这一步就相当于把每个Patch“编码”成模型更方便处理的向量。2.2 Patch Embedding用卷积实现竟然更快很多人第一次看到ViT代码时会愣一下不是说分块吗怎么用了Conv2dself.patch_embed nn.Conv2d( in_channels3, out_channelsembed_dim, kernel_sizepatch_size, stridepatch_size )这里一个kernel_size和stride都等于patch_size的卷积就能一步完成“切块 Flatten 线性映射”三个操作。你想想看卷积核在图像上滑动的步长等于核大小意味着每个卷积核的响应区域恰好是一个Patch互不重叠卷积核的输出通道数是embed_dim相当于对每个Patch做了一个D维的线性变换。这和“先切块、再拉平、再Linear”完全等价。用Conv2d实现的优势在于底层有im2col等优化而且GPU对卷积的加速非常成熟速度更快。实际工程中基本都用这个写法。2.3 Position Embedding让模型知道谁先谁后Self-Attention本身是对“集合”做运算它不考虑顺序。你把序列里的第1个token和第5个token互换位置Attention的计算结果在排列意义上是一样的。这在NLP里有问题词的顺序影响语义。在图像里同样有问题Patch的顺序就是空间布局信息打乱了图像就变成拼图了。解决办法是加上位置编码Position Embedding。ViT用的是可学习的位置编码也就是初始化一个[196, 768]的随机向量矩阵和Patch Embedding后的序列直接相加。这个矩阵在训练过程中会不断更新模型自动学习每个位置应该加什么样的“位置特征”。论文里对比过1D、2D位置编码也对比过可学习和固定编码最终效果差别不大ViT就用了最简单的一维可学习方式。2.4 CLS Token用一个特殊向量聚合全局信息BERT里有个[CLS] token它的设计初衷是序列经过Transformer编码后用这个特殊token对应的输出向量来做分类。ViT照搬了这个思路。具体做法是在196个Patch token前面拼接一个可学习的向量形状是[1, 768]这样序列长度变成197。这个CLS token在训练过程中会“吸收”整个序列的信息——因为它可以通过Attention关注到所有Patch。最后做分类时只取CLS token对应的输出接一个全连接层输出类别logits。为什么不直接对所有Patch的输出做平均池化因为平均池化是无差别的而CLS token是让模型自己决定“我要重点看哪些Patch”理论上表达力更强。虽然实验上两者差距不大但CLS token成了ViT系列的标准设计。2.5 Transformer EncoderAttention到底算了什么有了197个token接下来就是标准的Transformer Encoder Block。一个Block包含四个核心部分LayerNorm、Multi-Head Self-AttentionMSA、MLPFeed-Forward Network、残差连接。Self-Attention的公式长这样Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V对每个token通过三个不同的线性层算出Query、Key、Value。Query想表达“我想找什么”Key表达“我有什么”Value表达“我实际提供什么信息”。Query和所有Key做点积得到一个相关性分数除以sqrt(d_k)防止数值过大导致softmax饱和再过softmax变成权重最后用权重加权求和所有Value。用生活化的方式理解你在教室里找学习搭子。你心里想“我需要一个数学好的人”Query扫一眼全班同学每个人头上都贴着自己的“特长标签”Key你给数学好的同学打了高分Attention权重然后认真听他们讲题加权求和这些人的Value。Multi-Head就是把上面的过程重复H次每次用不同的QKV线性变换相当于从H个子空间分别找关系最后把结果拼接起来再过一层线性变换。多头的好处是不同的头可以关注不同的模式有的头关注相邻Patch的纹理连续性有的头关注远距离Patch的全局结构关系。每个Block的完整顺序是输入x先过LayerNorm再过MSA结果和x相加残差连接结果再过LayerNorm再过MLP结果再次和输入相加这个结构叫Pre-LN也就是先归一化再进子层。早期的Transformer是Post-LN先子层再归一化但Pre-LN训练更稳定特别是深层网络所以ViT沿用了Pre-LN。MLP部分通常是一个两层的全连接第一层把维度从D放大到4D激活函数用GELU第二层再映射回D。这个4D的放大比例和激活函数选择都是实践总结出的经验值直接沿用就好。2.6 为什么ViT需要大量数据ViT没有CNN那种“局部性”和“平移等变性”的先天假设所有空间关系全靠数据驱动学习。这意味着要让模型自己发现“相邻Patch通常相关性更高”这类规律需要足够多的样本来支撑。在ImageNet这种120万张图的数据集上从头训练的ViT效果不如同量级ResNet但数据量放大到几亿张时ViT反超。所以实际工程里使用ViT几乎都是加载预训练权重再微调很少从头训练。这个问题在后面代码实战部分会重点体现。3. 代码实现从零搭建一个Vision Transformer3.1 环境准备我用的是PyTorch 2.x版本Python 3.10以上torchvision用来加载数据集和做数据增强。代码不依赖timm等第三方库全部手写。pip install torch torchvision matplotlib tqdm3.2 Patch Embedding层import torch import torch.nn as nn class PatchEmbed(nn.Module): 把图像切成patch并映射到embed_dim维向量。 img_size: 输入图像边长默认224假设正方形 patch_size: patch边长默认16 in_chans: 输入通道数RGB为3 embed_dim: 映射后的向量维度 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, 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这里有一个常用技巧flatten(2)把H和W两维合并成序列维度transpose(1, 2)把序列维度调整到中间得到[B, num_patches, embed_dim]的形状符合Transformer的输入要求。3.3 Multi-Head Self-Attention实现class MultiHeadSelfAttention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasFalse, attn_drop0., proj_drop0.): super().__init__() assert dim % num_heads 0, dim必须能被num_heads整除 self.num_heads num_heads self.head_dim dim // num_heads self.scale self.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: [B, N, 3*C] qkv self.qkv(x) # 拆成Q、K、V并分成多头 qkv qkv.reshape(B, N, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # q, k, v: 每个都是 [B, num_heads, N, head_dim] q, k, v qkv.unbind(0) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.attn_drop(attn) x attn v # [B, num_heads, N, head_dim] x x.transpose(1, 2).reshape(B, N, C) x self.proj(x) x self.proj_drop(x) return x这里的高效写法是用一个Linear层同时算出Q、K、V再通过reshape和permute拆成多头避免写三个独立Linear的冗余计算。scale取head_dim的-0.5次方等价于除以sqrt(head_dim)。3.4 Transformer Encoder Blockclass MLP(nn.Module): def __init__(self, in_dim, hidden_dimNone, out_dimNone, act_layernn.GELU, drop0.): super().__init__() out_dim out_dim or in_dim hidden_dim hidden_dim or in_dim * 4 self.fc1 nn.Linear(in_dim, hidden_dim) self.act act_layer() self.fc2 nn.Linear(hidden_dim, out_dim) self.drop nn.Dropout(drop) def forward(self, x): x self.fc1(x) x self.act(x) x self.drop(x) x self.fc2(x) x self.drop(x) return x class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., qkv_biasFalse, drop0., attn_drop0., act_layernn.GELU): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn MultiHeadSelfAttention( dim, num_headsnum_heads, qkv_biasqkv_bias, attn_dropattn_drop, proj_dropdrop ) self.norm2 nn.LayerNorm(dim) self.mlp MLP(in_dimdim, hidden_dimint(dim * mlp_ratio), act_layeract_layer, dropdrop) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x注意看Pre-LN的写法x先经过norm1再过attn然后和原始x相加第二个子层同理。如果你写成attn(x)再norm就变成Post-LN了训练稳定性会差一截。3.5 完整ViT模型class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4., qkv_biasTrue, drop_rate0., attn_drop_rate0.): super().__init__() self.patch_embed PatchEmbed( img_sizeimg_size, patch_sizepatch_size, in_chansin_chans, embed_dimembed_dim ) num_patches self.patch_embed.num_patches # CLS token每个样本一个可学习向量 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) # 位置编码长度是 num_patches 1多出的1是CLS token的位置 self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(pdrop_rate) # 堆叠depth个Transformer Block self.blocks nn.Sequential(*[ TransformerBlock( dimembed_dim, num_headsnum_heads, mlp_ratiomlp_ratio, qkv_biasqkv_bias, dropdrop_rate, attn_dropattn_drop_rate ) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) # 初始化权重 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) self.apply(self._init_weights) def _init_weights(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.zeros_(m.bias) nn.init.ones_(m.weight) def forward(self, x): B x.shape[0] x self.patch_embed(x) # [B, 196, 768] # 把CLS token复制到每个样本并拼接到序列最前面 cls_tokens self.cls_token.expand(B, -1, -1) # [B, 1, 768] x torch.cat([cls_tokens, x], dim1) # [B, 197, 768] x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) # 只取CLS token对应的输出 cls_out x[:, 0] logits self.head(cls_out) return logits核心参数说明参数含义经典ViT-Base取值img_size输入图像边长224patch_sizePatch边长16embed_dimToken编码维度768depthEncoder Block数量12num_headsAttention头数12mlp_ratioMLP隐藏层放大倍数4num_classes分类类别数按任务定这里有个容易被忽略的细节为什么cls_token初始化为0而pos_embed用trunc_normal_CLS token初始化为0是让模型从头学这个特殊向量而位置编码的初始化通常设置成一个均匀的随机分布比较好trunc_normal_(std0.02)是实践中常用的初始化方法。3.6 验证模型能不能跑model VisionTransformer( img_size224, patch_size16, num_classes1000, embed_dim768, depth12, num_heads12 ) x torch.randn(4, 3, 224, 224) logits model(x) print(logits.shape) # torch.Size([4, 1000])一个Batch为4的张量顺利通过输出[4, 1000]。到这里一个完整可用的ViT就搭好了。3.7 在CIFAR-10上训练一个迷你版ViT直接用标准ViT-Base在CIFAR-10上训练不现实原因有两个一是模型122M参数CIFAR-10只有5万张训练图严重过拟合二是224x224输入在CIFAR-10上没有必要。实际实验可以把patch_size调小、embed_dim调小组成一个迷你ViT在保持思路不变的前提下让模型落到可训练规模。from torchvision import datasets, transforms from torch.utils.data import DataLoader import torch.optim as optim # 数据增强 transform_train transforms.Compose([ transforms.Resize(64), transforms.RandomCrop(64, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)) ]) transform_test transforms.Compose([ transforms.Resize(64), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)) ]) train_dataset datasets.CIFAR10(data, trainTrue, downloadTrue, transformtransform_train) test_dataset datasets.CIFAR10(data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4) test_loader DataLoader(test_dataset, batch_size256, shuffleFalse, num_workers4) # 迷你ViT model VisionTransformer( img_size64, patch_size8, num_classes10, embed_dim256, depth6, num_heads8 ) # 优化器AdamW是Transformer训练的标配 optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) criterion nn.CrossEntropyLoss() # 训练一轮的代码 model.train() for images, labels in train_loader: optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step()这个迷你模型参数量大约在700万左右在CIFAR-10上从零训练可以达到约70%左右的准确率。听上去不太高但这是ViT固有的数据饥饿特性没有预训练权重没有海量数据它学不过同等规模的CNN。想验证ViT的真实能力一定要用预训练权重做迁移学习。4. 实战中的坑与排查技巧4.1 位置编码维度报错这个坑几乎人人都会踩。如果你把模型训练好的权重复用到一个不同分辨率的输入上比如224x224训练的模型想改成384x384推理Patch数量从196变成576pos_embed的维度就对不上了。解决办法是插值对pos_embed做双线性插值把196扩展到576再微调几个epoch。4.2 小数据集过拟合严重ViT在小数据集上比CNN更容易过拟合。我的建议是加大数据增强RandomAugment、MixUp、CutMix都很有效增大Dropout比例attention层的dropout和MLP层的dropout都调到0.1以上减小patch_sizepatch从16改成8token数增加模型能捕获更多细节但也更吃显存尽早加载预训练权重不要从零训练4.3 Loss震荡不收敛如果训练时Loss在来回震荡先用小学习率跑几十步看看能不能下降。Transformer对学习率很敏感常用策略是warmup前几个epoch学习率从0线性升到设定值后面再按余弦或者线性衰减。AdamW的weight_decay一般设在0.05到0.3之间不要用SGD效果差很多。4.4 显存溢出ViT-Base在224输入、Batch128下大约需要14GB显存这还只是训练。如果显存不够优先降低Batch size其次减小输入分辨率。对图像分类任务来说小分辨率的损失没有想象中那么大很多实际项目直接用160x160或者192x192。4.5 常见问题速查表问题现象可能原因解决方法位置编码维度不匹配输入分辨率与预训练不一致对pos_embed做插值或重新训练Loss不下降学习率太高或太低使用warmup 更小的峰值学习率训练集准确率高但验证集低过拟合增大dropout、增强数据、减小模型混合精度训练报错某些Op在fp16下不稳定用autocast时对embedding层保持fp32推理结果全是同一类未加载预训练权重检查权重加载逻辑和类别对齐4.6 代码调试心得我建议你从最小配置开始调通patch_size8, depth2, num_heads4, embed_dim128先在单张图片上跑通前向和反向。确认无误后再逐步放大参数。这样即使出错也能快速定位是网络结构问题还是数据问题。5. 从ViT到变体了解技术演进方向ViT的意义不只是提供了一个图像分类模型它打开了一个思路把视觉问题转化成序列问题然后用Transformer统一处理。沿着这条路线业界快速演化出很多工作。DeiTData-efficient Image Transformers解决了ViT需要海量数据的问题通过知识蒸馏从CNN教师网络学习和更强数据增强在ImageNet上仅用120万张图就从零训练出好的ViT模型。Swin Transformer引入层级特征金字塔结构用Windows Self-Attention限制注意力计算范围再通过Shifted Window让信息跨窗口流动。它既能处理高分辨率输入又有较好的局部建模能力成了检测和分割任务的常用backbone。另外一些研究开始引入拓扑或高阶关系建模比如我近期关注的hgformerTopology-Aware Vision Transformer with Hypergraph Learning。它的切入点在于标准Attention建模的是两两Patch之间的关系但图像里可能存在“多个Patch共同构成一个语义区域”的高阶关系普通图结构表达不了这种多对多关系。hgformer引入了超图学习让超边Hyperedge连接多个Patch显式建模更高阶的语义关联。这类变体本质上没有脱离ViT的基础框架依然是Patch Embedding Transformer Encoder 任务头。区别在于Attention设计、token交互方式、层级结构这些模块的替换。因此把标准ViT的代码吃透再去看这些论文你会发现自己能快速定位“它改的是哪一块”而不会迷失在论文的公式里。在实际工程选型中我的建议是如果只是做标准图像分类且有预训练可用直接用预训练ViT或DeiT微调省时省力如果做检测分割选Swin这类层级化结构如果做细粒度识别或者数据量不大先考虑改数据增强和小模型不要一上来就上大模型。最后分享一个实操中的小技巧无论如何调整模型先跑一个Batch的过拟合测试——只喂几十张图看看模型能不能把训练集准确率冲到接近100%。如果连这一步都做不到说明网络结构有Bug不要急着调超参数。这是我在调试ViT过程中最有用的一个习惯。