Vision Transformer模型规格与源码实现全解析
1. 从“Transformer”到“Vision Transformer”一场视觉领域的范式革命如果你在2020年之前问我处理图像的主流模型是什么我会毫不犹豫地回答卷积神经网络CNN。从AlexNet到ResNet再到EfficientNetCNN架构在图像分类、目标检测等任务上建立了不可撼动的统治地位。然而2020年一篇名为《An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale》的论文彻底打破了这一局面。这篇论文提出的Vision TransformerViT将原本用于自然语言处理的Transformer架构直接搬到了计算机视觉领域其核心思想简单到令人惊讶把一张图片切成一个个小块Patch把这些小块当作句子中的“单词”Token来处理。这个想法初看有些“暴力”甚至有点反直觉。毕竟图像具有天然的二维空间结构而CNN的卷积操作天生就是为了捕捉这种局部相关性而设计的。Transformer的自注意力机制虽然能建模长距离依赖但它对序列的顺序不敏感需要额外引入位置编码。ViT的成功证明了当数据量足够大时一个足够强大的通用架构Transformer可以超越为特定领域图像精心设计的架构CNN。这不仅仅是技术上的突破更是一种思维范式的转变从“为视觉任务设计专用模型”转向“用通用架构处理视觉信号”。如今ViT及其变种已成为计算机视觉领域的基石模型理解其常见的模型规格和源码实现是深入现代视觉AI的必经之路。2. ViT模型规格全解析从Tiny到Huge的演进之路ViT论文中提出了几个标准化的模型规格这些规格主要区别在于Transformer编码器的深度层数和宽度隐藏层维度、注意力头数。理解这些规格的命名和参数是后续进行模型选择、微调乃至改进的基础。下面我将结合论文和后续社区实践详细拆解这些常见规格。2.1 标准ViT规格Base, Large, Huge最初的ViT论文主要聚焦于三个规格ViT-Base ViT-Large和ViT-Huge。它们的命名规则非常直观直接反映了模型的规模。ViT-Base (ViT-B/16)这是最常用、也是后续研究中最常作为基线的规格。“B/16”中的“16”指的是将图像分割成的每个小块的尺寸Patch Size是16x16像素。例如对于一张224x224的标准输入图像会被分割成 (224/16) * (224/16) 14 * 14 196个图像块。每个图像块经过线性投影后会变成一个768维的向量这就是Transformer的输入嵌入维度。层数 (Layers/Depth) 12层Transformer编码器。隐藏层维度 (Hidden Size, D) 768。这是每个图像块被编码后的向量维度也是Transformer内部前馈网络FFN的输入输出维度。多头注意力头数 (Heads) 12。每个注意力头的维度是 768 / 12 64。前馈网络中间层维度 (MLP Size) 通常为隐藏层维度的4倍即 768 * 4 3072。参数量 大约8600万86M参数。这个规模在当时的计算资源下已经属于“大模型”但相比后来的模型它算是轻量级的。ViT-Large (ViT-L/16)大型规格在更大规模的数据集如JFT-300M上预训练后展现了更强的性能。层数 24层。隐藏层维度 1024。注意力头数 16。MLP Size 1024 * 4 4096。参数量 大约3.07亿307M参数。性能相比Base有显著提升但计算开销也大幅增加。ViT-Huge (ViT-H/14)巨型规格注意这里的Patch Size变成了14。更小的Patch Size意味着更多的图像块对于224x224输入是16*16256个序列更长计算量急剧上升但也能捕捉更细粒度的信息。层数 32层。隐藏层维度 1280。注意力头数 16。MLP Size 1280 * 4 5120。参数量 超过6.32亿632M参数。这是典型的“大力出奇迹”的模型需要在海量数据和强大算力下才能充分训练。2.2 社区扩展规格Tiny, Small, Giant随着ViT的普及社区和后续工作如DeiT, Swin Transformer为了适应不同的计算预算和任务需求引入了更小或更特化的规格。ViT-Tiny (ViT-Ti/16) 和 ViT-Small (ViT-S/16)这两个规格在原始论文中没有但在Timm库、DeiT等工作中被广泛使用旨在提供更轻量、更快的选择便于在移动端或资源受限场景下部署。ViT-Tiny 通常为12层隐藏维度192注意力头数3参数量约600万。它牺牲了一些性能但速度极快。ViT-Small 通常为12层隐藏维度384注意力头数6参数量约2200万。它在速度和精度之间取得了很好的平衡是许多轻量级应用的起点。ViT-Giant 和 超大规模变体除了Huge一些研究探索了更大的模型如ViT-GGiant参数可能达到10亿甚至百亿级别。这些模型通常与特定的训练技巧如强大的正则化、更精细的优化策略和超大规模数据集绑定。注意模型规格的“灵活性”。以上参数是标准配置但在实际代码库如Hugging Face Transformers, Timm中你可以通过参数灵活调整。例如你可以创建一个“深度为8宽度为512”的自定义ViT。理解这些核心维度深度、宽度、头数、Patch大小的相互关系比死记硬背某个配置更重要。2.3 Patch Size的影响不仅仅是分辨率Patch Size是一个关键但容易被忽视的超参数。它直接决定了两个核心因素序列长度 序列长度 (图像高 / Patch Size) * (图像宽 / Patch Size)。序列长度直接影响自注意力机制的计算复杂度O(n²)是模型计算开销的主要决定因素之一。信息粒度 一个16x16的Patch会丢失很多细节而一个8x8甚至4x4的Patch能保留更多局部信息但代价是序列长度呈平方级增长。为什么ViT-H/14用14而不是16一个可能的原因是14能整除常见的图像尺寸如224同时相比16它略微增加了序列长度从196到256让模型在相同层数和宽度下有更多的“单词”可以处理可能有助于提升细节建模能力但这也显著增加了计算成本。在实际选择时需要根据任务对细节的需求和你的计算资源进行权衡。3. 源码深度剖析从图像块到分类结果只看论文和规格参数是远远不够的真正的理解来自于代码。下面我将以PyTorch风格结合Hugging Face Transformers库和原始论文的实现思路手把手拆解ViT的核心源码模块。我会解释每一行关键代码背后的意图而不仅仅是展示代码。3.1 图像分块与线性投影从2D到1D序列这是ViT的第一步也是最关键的一步——将2D图像转换为1D序列。import torch import torch.nn as nn class PatchEmbedding(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # 计算patch数量如14x14196 # 核心操作用一个卷积层同时完成“分块”和“投影” # 卷积核大小步长patch_size这样卷积核每次滑动刚好不重叠地覆盖一个patch # 输出通道数就是嵌入维度embed_dim self.projection nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) # 可学习的位置编码[1, num_patches 1, embed_dim] # 加1是因为还有一个额外的[class] token self.position_embeddings nn.Parameter(torch.zeros(1, self.num_patches 1, embed_dim)) # [class] token一个可学习的向量用于最终分类 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) def forward(self, x): # x形状: [batch_size, channels, height, width] batch_size x.shape[0] # 投影 [B, C, H, W] - [B, embed_dim, H/patch, W/patch] x self.projection(x) # 展平空间维度 [B, embed_dim, H/p, W/p] - [B, embed_dim, num_patches] x x.flatten(2) # 调整维度变成序列形式 [B, num_patches, embed_dim] x x.transpose(1, 2) # 扩展并添加[class] token cls_tokens self.cls_token.expand(batch_size, -1, -1) # [B, 1, embed_dim] x torch.cat((cls_tokens, x), dim1) # [B, num_patches1, embed_dim] # 加上位置编码 x x self.position_embeddings return x为什么用卷积层做投影这是实现上的一个巧妙之处。nn.Conv2d当kernel_sizestridepatch_size时它的效果等价于1将图像划分成不重叠的patch_size x patch_size的小块2对每个小块的所有像素patch_size*patch_size*channels个值进行一个线性变换全连接映射到embed_dim维空间。用卷积实现比手动循环切片再全连接高效得多也完全兼容GPU的并行计算。[class] token的奥秘这个特殊的token被预先添加到序列开头。它不包含任何具体的图像块信息但在经过所有Transformer层后它“汇聚”了整个序列的全局上下文信息。最终我们只用这个cls_token的输出向量送入分类头进行图像分类。这比对所有图像块的特征做平均池化另一种做法更灵活因为它允许模型学习如何聚合信息。3.2 Transformer编码器自注意力与前馈网络的核心这是ViT的“发动机”由多个相同的层堆叠而成。每一层主要包含多头自注意力MSA和多层感知机MLP两个子层每个子层前后都有层归一化LayerNorm和残差连接Residual Connection。class TransformerEncoderLayer(nn.Module): def __init__(self, embed_dim768, num_heads12, mlp_ratio4.0, dropout0.0): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attention nn.MultiheadAttention(embed_dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(embed_dim) self.mlp nn.Sequential( nn.Linear(embed_dim, int(embed_dim * mlp_ratio)), nn.GELU(), # ViT中使用GELU激活函数比ReLU更平滑 nn.Dropout(dropout), nn.Linear(int(embed_dim * mlp_ratio), embed_dim), nn.Dropout(dropout) ) def forward(self, x): # 残差连接1: 注意力子层 # 先做层归一化再计算注意力 x_norm1 self.norm1(x) attn_output, _ self.attention(x_norm1, x_norm1, x_norm1) x x attn_output # 残差连接 # 残差连接2: MLP子层 x_norm2 self.norm2(x) mlp_output self.mlp(x_norm2) x x mlp_output # 残差连接 return xPre-Norm vs Post-Norm注意看代码这里采用的是Pre-Norm结构先做LayerNorm再进行注意力或MLP计算。这是原始Transformer论文《Attention is All You Need》中Post-Norm结构的一种变体。大量实践包括ViT、GPT等表明Pre-Norm结构在训练深度网络时更加稳定梯度更容易回传是当前的主流选择。如果你看到代码里是x x self.attention(self.norm1(x))那就是Pre-Norm。GELU激活函数ViT中MLP使用的是GELU高斯误差线性单元而不是CNN中更常见的ReLU。GELU可以看作是ReLU的一个平滑随机版本它在负值区域也有微小的输出理论上能提供更丰富的梯度信息在实践中对深层Transformer的优化有积极影响。3.3 位置编码为序列注入空间信息Transformer本身是置换不变的Permutation-Invariant打乱输入序列的顺序输出序列的顺序也会被打乱但内容不变。这对于图像来说是不可接受的因为图像的“左”和“右”有明确的语义。因此必须显式地加入位置信息。ViT使用的是可学习的一维位置编码。如PatchEmbedding代码所示self.position_embeddings是一个可学习的参数形状为[1, num_patches1, embed_dim]。每个位置包括[class] token的位置都有一个独一无二的、可学习的向量。模型在训练过程中会学会如何利用这些向量来理解图像块之间的空间关系。注意位置编码的争议与演进。可学习的位置编码简单有效但缺乏对相对位置关系的明确归纳偏置。后续很多工作探索了其他形式如相对位置编码Swin Transformer、二维正弦编码、甚至条件位置编码CPVT。在源码阅读时留意position_embeddings的初始化和使用方式是理解模型空间感知能力的关键。3.4 整体ViT模型组装将以上所有部分组合起来再加上一个分类头就构成了完整的ViT模型。class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0): super().__init__() self.patch_embed PatchEmbedding(img_size, patch_size, in_channels, embed_dim) self.encoder_layers nn.ModuleList([ TransformerEncoderLayer(embed_dim, num_heads, mlp_ratio) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) # 最终输出前的层归一化 self.head nn.Linear(embed_dim, num_classes) # 分类头 def forward(self, x): # 1. 图像分块与嵌入 x self.patch_embed(x) # [B, num_patches1, embed_dim] # 2. 通过所有Transformer编码层 for layer in self.encoder_layers: x layer(x) # 3. 取[class] token对应的输出并做最终归一化 x self.norm(x) # 对整个序列做归一化 cls_token_final x[:, 0] # 取出第一个token即[class] token # 4. 分类 logits self.head(cls_token_final) return logits前向传播流程梳理输入图像[B, 3, 224, 224]。经过PatchEmbedding得到[B, 197, 768]196个图像块 1个[class] token。经过12层以ViT-B为例TransformerEncoderLayer形状保持不变。对最终输出序列进行层归一化。提取序列中第一个向量即[class] token对应的输出形状为[B, 768]。通过一个线性分类头映射到类别数如1000得到最终的分类逻辑值[B, 1000]。4. 关键源码细节与实战避坑指南阅读标准实现只是第一步在实际使用、修改或调试ViT时有几个细节至关重要也是容易踩坑的地方。4.1 分类头与预训练权重的适配当你从Hugging Face或Timm加载一个在ImageNet-21k上预训练的ViT-B/16模型并想用它做10分类的任务时直接使用会报错。因为预训练模型的分类头head输出维度是21843ImageNet-21k的类别数而你的任务需要10。正确做法是替换分类头from transformers import ViTForImageClassification model ViTForImageClassification.from_pretrained(google/vit-base-patch16-224-in21k) # 替换分类器 model.classifier nn.Linear(model.config.hidden_size, 10) # 假设你的任务有10类或者在使用Timm库时import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10) # timm的create_model函数会帮你自动处理分类头的适配非常方便。避坑点不要仅仅冻结前面的层然后训练新分类头。对于Transformer通常建议进行整体微调。因为自注意力机制是全局的即使是为ImageNet-21k预训练的底层特征也可能需要根据你的新数据分布进行细微调整。你可以使用较低的学习率如预训练层的1/10来微调整个模型这通常比只训练分类头效果更好。4.2 位置编码的可视化与理解理解模型学到了什么位置信息的一个好方法是可视化学习到的位置编码。我们可以取出position_embeddings参数然后用余弦相似度计算不同位置编码之间的相关性。import matplotlib.pyplot as plt import numpy as np # 假设 model 是你的ViT模型 pos_embed model.patch_embed.position_embeddings[0, 1:].detach().cpu().numpy() # 去掉[class] token的位置编码 num_patches pos_embed.shape[0] h w int(np.sqrt(num_patches)) # 假设是正方形排列 # 计算所有位置编码之间的余弦相似度矩阵 similarity_matrix np.zeros((num_patches, num_patches)) for i in range(num_patches): for j in range(num_patches): similarity_matrix[i, j] np.dot(pos_embed[i], pos_embed[j]) / (np.linalg.norm(pos_embed[i]) * np.linalg.norm(pos_embed[j])) # 可视化 plt.figure(figsize(10, 8)) plt.imshow(similarity_matrix, cmaphot, interpolationnearest) plt.colorbar() plt.title(Cosine Similarity between Learned Position Embeddings) plt.xlabel(Patch Index) plt.ylabel(Patch Index) plt.show()如果模型学到了合理的空间结构你可能会看到相似度矩阵呈现出“块状”结构相邻位置如图像中上下左右的编码相似度更高。这可以作为模型是否正常工作的一个辅助检查。4.3 注意力权重的分析与可视化自注意力机制是Transformer的灵魂。可视化注意力图可以帮助我们理解模型在做出决策时“看”向了图像的哪些部分。# 这是一个简化的示例实际中需要修改模型forward以返回注意力权重 def forward_with_attention(self, x): x self.patch_embed(x) attentions [] for layer in self.encoder_layers: # 假设我们在attention层保存了权重 x, attn_weights layer.attention(self.norm1(x), self.norm1(x), self.norm1(x), need_weightsTrue) attentions.append(attn_weights.detach()) # ... 继续前向传播 return x, attentions # 获取最后一层某个头例如头0的注意力图关注[class] token对其他所有patch的注意力 # attn_weights形状: [batch, num_heads, seq_len, seq_len] cls_attentions attentions[-1][0, 0, 0, 1:] # 取batch0, head0, [class] token对所有图像块的注意力 cls_attentions cls_attentions.reshape(h, w) # 重塑为二维网格 plt.figure(figsize(8, 8)) plt.imshow(cls_attentions, cmapReds) plt.title(Attention from [CLS] token to patches (Last Layer, Head 0)) plt.axis(off) plt.show()通常我们会发现浅层的注意力比较分散关注局部纹理而深层的注意力更加集中和语义化可能会聚焦于物体的关键部位如狗的头部、车轮。这直观地展示了ViT如何从局部到全局整合信息。4.4 混合架构当CNN遇见ViT纯粹的ViT在中小型数据集上训练容易过拟合因为它缺乏CNN固有的平移不变性和局部性归纳偏置。一个有效的技巧是使用混合架构Hybrid Architecture即用CNN的骨干网络如ResNet来提取图像特征图然后将这个特征图而不是原始图像划分成块送入Transformer。在源码层面这只需要修改PatchEmbedding模块。不再用一个大卷积核做投影而是用一个轻量的CNN如几个卷积层来生成一个低分辨率、高通道数的特征图然后将这个特征图的每个空间位置视为一个“Patch”。class HybridPatchEmbedding(nn.Module): def __init__(self, backbone, feature_map_channels, embed_dim): super().__init__() self.backbone backbone # 一个预训练或随机初始化的CNN self.feature_proj nn.Linear(feature_map_channels, embed_dim) # 将CNN通道数投影到Transformer维度 def forward(self, x): # 通过CNN骨干网络 with torch.no_grad(): # 微调时可以选择冻结CNN features self.backbone(x) # 假设输出为 [B, C, H, W] # 将特征图视为序列 B, C, H, W features.shape features features.flatten(2).transpose(1, 2) # [B, H*W, C] # 投影到embed_dim x self.feature_proj(features) # [B, num_patches, embed_dim] # 后续添加[class] token和位置编码的步骤与标准ViT相同 # ... return x这种混合架构结合了CNN在早期特征提取上的优势和Transformer在长距离建模上的能力在数据量有限时往往能取得比纯ViT更好的效果。在Timm库中你可以通过模型名如vit_base_resnet50_224来使用这种混合模型。