ViT图像分类实战:基于Pytorch从零实现Vision Transformer

📅 发布时间:2026/9/1 2:28:41
ViT图像分类实战:基于Pytorch从零实现Vision Transformer 简介VisionTransformer-Pytorch 是一份基于 PyTorch 的视觉 Transformer 实现资源面向有一定深度学习基础、希望快速在图像分类任务中应用 ViT 模型的开发者。项目采用类似 EfficientNet 的简洁 API 设计通过 pip 安装后即可加载 ViT-B_16 预训练权重省去自行搭建与训练的时间。压缩包共 14 个文件以 9 个 Python 源文件为主涵盖核心模型、ResNet 辅助模块、图像处理工具、JAX 权重转换脚本与测试用例另有 README 与许可证文件说明使用方式及授权信息整体仅 25KB轻量易读。资源包含从 JAX 预训练权重迁移至 PyTorch 的转换工具便于获得官方参数同时提供示例脚本与测试代码帮助理解 ViT 的输入输出细节。目前已有 1163 人学习适合作为 ViT 实现参考或快速集成到视觉任务的起点。 我最近把一个图像分类项目全部重写了一遍核心模型从ResNet换成了VisionTransformer代码全部基于Pytorch实现。折腾了大概一周时间把ViT从零到一完整搭了出来包括数据管线、模型结构、训练脚本、推理加速整个过程踩了不少坑也积累了一些实际经验。这篇就把这个项目的完整思路和实现细节拆开讲讲给正在入门或准备上手ViT的朋友一个可以直接参考的实操方案。这个项目最核心的事情是把VisionTransformer简称ViT用Pytorch从底层代码开始实现并在实际数据集上完成训练和验证。它能解决的核心问题是图像分类任务中的特征提取和长距离依赖建模区别于传统CNN的局部感受野机制ViT通过自注意力机制能直接建模图像任意两个位置之间的依赖关系。适合对Transformer架构有了解、想把它迁移到视觉任务上的读者也适合刚学完Pytorch基础、想通过一个完整项目串联起数据加载、模型搭建、训练验证全流程的朋友。1. ViT的架构拆解为什么要把图像变成序列1.1 核心思路图像分块与序列化VisionTransformer的核心操作其实很简单把一张图片切成固定大小的patch然后把每个patch当作一个“词向量”送入Transformer编码器。我以ViT-B/16为例说明输入图片尺寸是224x224patch size设为16x16那么图片会被切分成(224/16)² 196个patch加上一个用于分类的[CLS] token序列长度就是197。这里有个非常关键的细节每个patch在送入Transformer之前需要经过一个线性映射把16x16x3RGB三通道的原始像素展平成一个768维的向量。我第一遍实现时很自然就想用nn.Linear来做这个映射但实际写的时候发现直接用nn.Conv2d更高效——用一个in_channels3, out_channels768, kernel_size16, stride16的卷积一步就把切分和线性投影同时完成了。这在算力上比先切patch再逐个做线性变换要省很多也是很多开源实现里普遍采用的做法。1.2 位置编码与[CLS] Token的设计逻辑序列化之后遇到的一个问题是Transformer本身是置换不变的它不知道序列里哪个patch在前哪个在后。为了让模型能感知patch的空间位置必须给每个token加上位置编码。ViT原文使用的是可学习的1D位置编码直接初始化一个nn.Parameter形状是(197, 768)和token embedding直接相加。[CLS] token的设计是从BERT里继承过来的思路在序列最前面拼接一个特殊的可学习向量经过Transformer编码器后用这个位置对应的输出向量作为整张图片的全局特征表示再接到分类头上。为什么不直接用所有token的平均池化我试验过两者的差异[CLS] token在训练中会主动聚合全图信息收敛速度和最终精度都比平均池化稍好一点这可能是因为它在注意力层中承担了类似“信息汇聚中心”的角色梯度信号也更集中。补一个非常容易踩的坑位置编码是加到序列维度上的一定要在patch embedding之后、送入Encoder之前加。我第二版代码里把位置编码加错了位置导致训练loss怎么都降不下去排查了很久才发现是维度broadcast出了问题。2. 环境搭建与Pytorch安装要点2.1 创建干净的运行环境动手写代码之前先把环境搞定这一步看似基础但最容易出问题。我的建议是使用conda创建一个独立的虚拟环境不要直接装在base环境里不然后面不同项目之间的依赖冲突会让人非常崩溃。创建命令很简单conda create -n vit python3.10 conda activate vitPython版本我建议选3.9或3.10这两个版本对Pytorch各版本的支持最稳妥。太新的Python版本比如3.13刚出来时可能会导致部分依赖库还没有对应的wheel包装起来很麻烦。2.2 GPU版Pytorch的安装细节GPU版本的安装是很多人卡壳的地方。核心原则是先确认你的CUDA版本再安装对应的Pytorch版本顺序不能反。查看CUDA版本用nvidia-smi如果显示的是CUDA 11.8那么对应的安装命令是pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118如果国内下载速度太慢我建议优先使用清华镜像源或者设置pip超时时间加长pip install torch torchvision torchaudio -i https://pypi.tuna.tsinghua.edu.cn/simple这里有个实际经验很多人会纠结到底下载哪个CUDA版本其实只要满足Pytorch要求版本 驱动支持的CUDA版本这个条件就可以。比如驱动支持CUDA 12.4那么你装cu118、cu121、cu124对应的Pytorch都能跑起来Pytorch会通过自带的CUDA runtime运行不依赖系统全局的CUDA。我当初装的时候直接选了大版本匹配的cu118版本实测跑ViT-B/16训练和推理都没问题。注意装完Pytorch后一定要验证GPU是否可用用python -c import torch; print(torch.cuda.is_available())输出True才说明环境没问题。很多人在这步输出False通常是Pytorch装成了CPU版本卸载重装GPU版即可。3. 核心模块的代码实现与参数解读3.1 Patch Embedding层直接看代码先把patch embedding层的实现贴出来import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels3, patch_size16, embed_dim768): super().__init__() self.patch_size patch_size self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (batch_size, 3, 224, 224) x self.proj(x) # (batch_size, embed_dim, 14, 14) x x.flatten(2) # (batch_size, embed_dim, 196) x x.transpose(1, 2) # (batch_size, 196, embed_dim) return x这里flatten(2)是把(14, 14)两个空间维度展平成196然后transpose(1, 2)把序列维度放到中间得到(batch, seq_len, embed_dim)这个形状是Transformer Encoder的标准输入格式。我第一版代码写的是x x.view(x.size(0), x.size(1), -1)效果和flatten(2)一样但flatten更直观后续维护也更清晰。3.2 Multi-Head Self-Attention实现自注意力是ViT的核心计算模块。它的思想可以这样理解对序列里的每个token通过三个不同的线性投影分别得到Query、Key、Value三个向量然后计算Query和所有Key的点积得到注意力权重再用softmax归一化最后用权重加权求和Value。多头注意力就是把embedding维度切成多份每个头独立做注意力计算最后拼接起来。class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim768, num_heads12, dropout0.1): super().__init__() self.num_heads num_heads self.head_dim embed_dim // num_heads self.qkv nn.Linear(embed_dim, embed_dim * 3) self.attn_drop nn.Dropout(dropout) self.proj nn.Linear(embed_dim, embed_dim) def forward(self, x): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn attn.softmax(dim-1) attn self.attn_drop(attn) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) return x这里有个关键的缩放因子self.head_dim ** -0.5很多新手会忽略它。这个缩放是为了防止点积结果过大导致softmax进入饱和区、梯度趋近于零。我在第一次实现时忘了加这个缩放训练时发现attention分布很快就变得非常尖锐模型几乎不更新了加上之后训练立刻变得稳定。3.3 Transformer Encoder与ViT主体一个完整的Encoder Layer由自注意力、MLP和两个LayerNorm组成采用了残差连接的思路。MLP的结构是两层全连接中间用GELU激活函数class TransformerBlock(nn.Module): def __init__(self, embed_dim768, num_heads12, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn MultiHeadSelfAttention(embed_dim, num_heads, dropout) self.norm2 nn.LayerNorm(embed_dim) hidden_dim int(embed_dim * mlp_ratio) self.mlp nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_dim, embed_dim), nn.Dropout(dropout) ) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x注意这里的顺序是x attn(norm(x))也就是Pre-LN结构先归一化再进注意力模块。这和原始Transformer的Post-LN先注意力再归一化有区别。Pre-LN在训练中更稳定能支持更大的学习率和更深的网络深度ViT官方实现用的也是这个结构我建议照做。最后把整个ViT主体组装起来PatchEmbed、CLS token、位置编码、多个TransformerBlock按顺序堆叠、LayerNorm和分类头。我这里用num_layers12配置了12层Encoder这也是ViT-Base的标准深度。4. 训练配置、数据准备与迁移学习技巧4.1 数据预处理与增强策略ViT在小数据集上比CNN更容易过拟合所以数据增强非常关键。我用的CIFAR-10数据集本身只有5万张训练图如果不做增强ViT-Base的参数量8600万远远超过数据量模型很快就在训练集上跑到99%准确率测试集却只有70%出头。我采用的增强策略是随机水平翻转、随机裁剪crop到224x224、RandAugment一个自动数据增强方法随机组合多种图像变换。实测下来RandAugment的提升最大能提高约3-4个百分点的测试精度。这里有个容易忽略的地方验证集的预处理不需要增强但需要做和训练集相同的resize和归一化操作归一化的均值和方差要提前统计数据集得到直接用ImageNet的默认值也能跑但精度会略差一点。4.2 训练超参设置ViT的优化配置和CNN有差异我从实际项目中验证过的参数配置如下超参数数值说明optimizerAdamW比Adam多了权重衰减解耦更适合Transformerbase learning rate3e-4过大会导致attention训练不稳定weight decay0.05ViT对weight decay比较敏感batch size256多卡可扩大到1024需同步调大lrwarmup epochs10让学习率从0线性升到目标值total epochs100足够观察到收敛趋势lr schedulecosine decay配合warmup效果最好这里重点说一下AdamW和warmup。AdamW把权重衰减从L2正则里解耦出来直接作用在参数更新上对Transformer架构有更好的泛化效果这在ViT的训练中几乎是标配。Warmup则是让模型在训练初期用很小的学习率先“稳住”避免注意力矩阵在初始化阶段就被更新到不理想的方向我的经验是10个epoch的warmup比较合理太短起不到作用太长会拖慢收敛。4.3 模型冻结与迁移学习很多实用场景下我们不一定要从头训练ViT。我建议直接加载在ImageNet上预训练好的权重然后在自己的小数据集上做微调这样能把训练时间从几天缩短到几小时。加载预训练权重的代码很简单可以从Pytorch官方或timm库获取import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue) model.head nn.Linear(768, num_classes)微调时有个很有用的技巧先把整个模型冻结requires_gradFalse只训练分类头和位置编码因为新任务输入尺寸可能变化位置编码需要重新插值跑几个epoch之后再把全部参数解冻进行全量微调。这样分阶段训练的好处是前期让新初始化的分类头先适应特征提取器后面解冻时不会因为分类头的随机梯度太大而破坏预训练特征。我在CIFAR-10上按这个方法微调100个epoch从头训练能达到94%左右而迁移学习只训练30个epoch就能到96%差距非常明显。5. 踩坑记录与问题排查技巧5.1 训练不收敛损失震荡不下降这是最常见的问题我之前也遇到过。排查思路按顺序来第一检查数据归一化是否正确看看输入图片的像素范围是否和数据增强的预期值一致。第二确认学习率是否过大ViT对学习率比CNN敏感得多CNN常用的1e-3学习率在ViT上很可能直接导致loss爆炸。第三检查warmup是否开启不使用warmup时前几个iteration的loss可能会出现很大的尖峰。第四确认attention中是否加了缩放因子这个前面提过漏加会导致梯度消失。5.2 显存不足OOM的优化手段ViT在训练时对显存的需求比同等级的ResNet大不少。如果你用的是12GB显存在batch size和环境允许范围内还是可能OOM可以从几个方向优化降低batch size配合梯度累积模拟大的batch开启混合精度训练AMP能显著减少显存占用代码上只需要加几行scaler torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, labels)使用gradient checkpointing以少量计算换显存torch.utils.checkpoint可以做到。混合精度是最推荐的做法ViT中的大部分操作对低精度并不敏感实测混合精度下训练速度能提升30%以上显存占用几乎减半。5.3 Patch Size与位置编码的联动影响这个坑比较隐蔽。当你要把ViT从一个输入尺寸迁移到另一个尺寸时位置编码需要做插值。比如预训练模型是224x224你想在384x384的输入上微调patch数量会从196变成576原始的位置编码无法直接使用需要做2D插值import torch.nn.functional as F pos_embed model.pos_embed pos_embed_interp F.interpolate( pos_embed[:, 1:, :].reshape(1, 14, 14, 768).permute(0, 3, 1, 2), size(24, 24), modebicubic ).permute(0, 2, 3, 1).reshape(1, 576, 768) new_pos_embed torch.cat([pos_embed[:, :1, :], pos_embed_interp], dim1)这里有个细节插值时[CLS] token的位置编码不需要插值直接保留原值只对patch位置部分做resize。如果这个操作做错了模型精度会大幅下降甚至可能直接不收敛。5.4 过拟合识别的三个信号如何判断模型是否过拟合我总结三个信号训练loss持续下降但验证loss开始回升训练精度接近100%而验证精度停滞attention可视化中模型过度关注某个局部区域。应对策略有两个一个是增强正则化强度dropout、weight decay、更多增强另一个是换更小规模的模型变体如ViT-Small。这里附一个不同规模ViT的具体参数量对比方便你根据数据集大小选型模型层数隐藏维度注意力头数参数量ViT-Tiny121923570万ViT-Small1238462200万ViT-Base12768128600万ViT-Large2410241630700万我个人的经验是小于10万张图的数据集优先考虑ViT-Tiny或ViT-Small配合预训练微调数据集在10万级以上才考虑ViT-Base及以上规模否则过拟合的风险非常高。一个小建议可视化Attention是理解ViT的好方法写代码之余我强烈建议你花点时间把attention map可视化出来。做法很简单在forward的时候保存最后一层的attention权重然后把每个head的注意力图resize回原始图片大小叠加到原图上。你会很直观地看到模型在关注哪些区域也能从中判断模型学到的特征是全局的还是碎片化的。我在调优过程中靠这个工具发现了不少问题比如某个head的attention完全塌缩成平均分布后来发现是dropout设置太高导致信息丢失。这个项目做下来我对ViT的理解比只看论文要深刻得多。从PatchEmbed到Self-Attention每一个模块的实现细节都值得反复推敲。如果大家照着实现过程中遇到任何问题欢迎在评论区交流我尽量抽空回复。本文还有配套的精品资源点击获取