尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Vision Transformer详解:从原理到PyTorch图像分类实战

Vision Transformer详解:从原理到PyTorch图像分类实战 简介VIT(vision transformer)实现图像分类完整项目包面向需要掌握Transformer在计算机视觉中应用的开发者、研究者和高年级学生也适用于深度学习课程设计、毕业设计或相关课题的复现起点可覆盖图像识别、场景分类等基础任务有助于建立从理论到落地的完整认知。它将图像切分为patch序列利用自注意力机制提取全局特征代表Transformer首次成功应用于CV领域与CNN局部卷积形成思路互补适合具备Python和基础深度学习知识的人群学习。压缩包约539.35MB约有2000个文件核心内容包括大批jpg图片样本、Python源代码和训练好的pth权重同时配有XML标注、JSON配置、类别名称等辅助文件目录规划清晰便于按需求查阅。目前该资源已有12407人学习浏览。随包提供可直接运行的数据集、完整训练与推理源码以及分类精度超过99%的预训练权重基本可以免去数据准备和环境适配成本代码覆盖数据读取、ViT网络构建、训练及测试全流程读者可据此快速复现基线也可替换数据集、调整超参数或加入可视化模块进一步理解注意力机制和图像分类建模思路。项目内置完整数据与权重省略数据收集和长时训练环节方便快速跑通并产出对比结果。1. 图像分类任务往 Transformer 上搬第一步是先把像素排列成序列一张 224×224 的彩色图片在 CNN 眼里是一个三维张量在 ViT 眼里则是 196 个长度为 768 的向量排成的序列。Vision Transformer 把图像切成固定大小的 patch每个 patch 线性映射成一个 token再交给标准 Transformer encoder 处理。分类结果不是来自某个全连接层直接压平特征图而是来自序列中一个专门训练的 class token 对应的输出向量。这个改动听起来不复杂但它确实打破了卷积在视觉任务里的垄断地位。ViT 没有卷积的局部归纳偏置它用全局自注意力建模任意两块区域之间的关系。因此在数据量足够大时ViT 在大图、细粒度任务上的上限往往高于同参数量的 CNN数据量不够时它又比 CNN 更容易过拟合这是个绕不开的权衡。读者如果是刚开始接触 vision transformer或者已经用 CNN 跑过几个图像分类项目、想换个结构看效果这篇文章能帮你把 ViT 的每个零件拆开跑通一个可复现的 PyTorch 实现再把训练参数和常见坑一次说清。2. ViT 的三个基础构件Patch Embedding、位置编码与 class tokenViT 出现之前Transformer 处理的都是离散 token文本天然就是序列。要让 Transformer 处理图像得先把连续像素变成它读得懂的向量序列。整个过程可以拆成三步切 patch 并投影成 embedding加位置编码在序列前面放一个 class token用它的输出去接分类头。这三件事都不难但每件都有值得注意的细节。2.1 Patch Embedding 做了什么为什么用卷积实现也行ViT 的输入是一个形状为 H×W×C 的图像张量。以 224×224 的 RGB 图为标准输入如果把 patch 大小定为 16单边能分出 14 份总共得到 14×14196 个 patch。每个 patch 在空间上是 16×16通道数是 3拉平后就是 768 维向量正好对应 ViT-Base 的隐藏维度。import torch import torch.nn as nn def patchify(x, patch_size): B, C, H, W x.shape num_h H // patch_size num_w W // patch_size x x.reshape(B, C, num_h, patch_size, num_w, patch_size) x x.permute(0, 2, 4, 3, 5, 1) # [B, num_h, num_w, patch_h, patch_w, C] x x.reshape(B, num_h * num_w, patch_size * patch_size * C) return x这个函数把 224×224×3 的图先切成 14×14 个子块再把每个 patch 内的 16×16×3 个像素拉平成 768 维向量。permute 和 reshape 的顺序有讲究先 permute 把空间的两个维度提到 patch 维之后再做 reshape 才能保证每个向量确实来自同一个 patch否则像素会串块。实际工程里很少手写这套 reshape。标准做法是用一个 stride 等于 kernel size 的卷积层替代——nn.Conv2d(in_chans3, embed_dim768, kernel_size16, stride16)一步就能完成切块和线性映射。卷积在这里没有充当特征提取器它只是把滑动窗口内的像素加权求和等价于对每个 patch 做全连接。用卷积实现的好处是前向速度快、PyTorch 内部对 Conv2d 的融合优化比逐 patch 的 einsum 更好。2.2 位置编码让 Transformer 知道序列里谁在前谁在后自注意力机制本身是对称的把序列打乱顺序Attention 计算出的权重完全不变。图像 patch 的顺序代表了空间结构丢掉它等于把一张图变成一堆互不相干的小碎片。于是位置编码是必需的。ViT 采用可学习位置编码不是 NLP 里常见的正弦函数。它是一个形状为 (1, num_patches 1, embed_dim) 的参数张量直接加到 patch embedding 的结果上。加号后面的 1 来自下一节要说的 class token。这个张量通常用截断正态分布初始化标准差设为 0.02而不是用零初始化否则初始阶段模型对位置的感知完全对称训练早期会多走弯路。可学习编码天然带有一个假设图片分辨率固定。输入从 224 变成 448patch 数就变成 28×28784原来的位置编码参数没法直接用。常见做法有二一是调整图片尺寸让它保持在 patch size 的整数倍二是做二维插值把位置编码张量从 (1, 197, 768) 插值到 (1, 785, 768)。第二种在迁移学习中偶有遇到但效果不如直接在目标分辨率上预训练来得稳妥。2.3 class token 承担分类的这个设计细节标准 CNN 分类是在最后接一个全局平均池化再连一层全连接。ViT 的原始论文没有走这条路而是在 patch 序列前面额外拼一个初始化为零的向量它不包含任何图像信息唯一的任务是在自注意力机制里不断和所有 patch 信息做交互最后从它对应的输出向量做分类。用 class token 的好处有两个。其一它给模型一个显式的信息汇聚点所有 patch 都会向它传递特征如果去掉它改用全局平均池化模型仍然能训练但 Attention 层没有一个明确的目标 token收敛速度通常稍慢。其二class token 的存在让 ViT 更贴近 BERT 的架构形式做迁移学习和多任务时逻辑更统一。需要注意class token 虽然初始是零向量但只要经过一次自注意力就会带上全局信息所以它不需要像手工特征那样设计。整个序列进入 encoder 前class token 的位置编码是单独一排参数与 patch 共用同一个位置编码张量。代码实现时记得把 patch embedding 和 class token 在序列维度拼在一起后再加位置编码。3. 用 PyTorch 写一个能跑的 ViT 图像分类模型这一章给出一个可以在单张消费级显卡上跑通的最小实现。代码不追求和官方实现逐行一致但结构完整涵盖了 Patch Embedding、Transformer encoder、分类头三大部分。先确定配置表再贴模型代码最后是训练主循环。3.1 数据准备与参数配置以 CIFAR-10 为例输入尺寸需要先适配 patch size 的整数倍。CIFAR-10 原始分辨率是 32×32如果 patch_size4单边分成 8 份共 64 个 patchembed_dim 可以适度缩小。这里按一个常见的“小 ViT”配置来跑参数值说明输入尺寸32×32原始分辨率不被裁切patch_size432 能被 4 整除边界干净隐藏维度192缩小版 embed_dim显存压力小深度6encoder 层数注意力头数3192/643头维度保持 64类别数10CIFAR-10优化器AdamW需设置 weight_decay批次大小128单卡可跑import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_set datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) test_set datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_set, batch_size128, shuffleTrue, num_workers4) test_loader DataLoader(test_set, batch_size256, shuffleFalse, num_workers4)数据增强先只保留随机裁剪和水平翻转这是最基本的组合。随机裁切的本质是让模型在每个 epoch 看到不同位置的图像内容变相扩大训练样本量。Normalize 的均值和标准差是 CIFAR-10 数据集的官方统计值迁移到别的数据集就换上对应统计量。使用 num_workers 时注意在 Windows 上可能遇到多进程加载问题改成 0 或 2 即可。3.2 ViT 模型定义模型拆成三个文件块更清晰PatchEmbed、Attention、TransformerBlock最后组装成 ViT 主类。import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim192): super().__init__() 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 self.proj(x) # [B, embed_dim, 8, 8] x x.flatten(2) # [B, embed_dim, 64] x x.transpose(1, 2) # [B, 64, embed_dim] return xprogress 将 224×224 的图转换成 (B, 196, 768) 的四维中间结果这是后面进入 encoder 的标准格式。flatten 发生在通道维度的后面先保留空间位置顺序再 transpose 把序列维度放到第二位和前文 patchify 函数的手工做法结果一致。接下来是自注意力和编码块。为了便于展示维度流动自注意力采用手写 qkv 的方式class Attention(nn.Module): def __init__(self, dim, num_heads3): super().__init__() self.num_heads num_heads self.head_dim dim // num_heads self.qkv nn.Linear(dim, dim * 3) self.proj nn.Linear(dim, 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) # [3, B, heads, N, head_dim] q, k, v qkv.unbind(0) attn (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn attn.softmax(dim-1) x (attn v).transpose(1, 2).reshape(B, N, C) return self.proj(x) class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn Attention(dim, num_heads) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(dropout), ) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x注意力计算里的缩放因子是 head_dim 的平方根倒数。如果不做缩放矩阵内积的方差会随维度增大而增大softmax 会过早进入饱和区梯度变得极小。标准 Transformer 里 dim768、heads12 时head_dim64缩放因子为 1/8。这里配置为 192/364保持同样的比例。残差连接放在 attention 和 mlp 外面。每个子层先做 LayerNorm 再进入注意力这个顺序来自 Pre-LN 结构能够大大缓解深层 Transformer 的梯度消失问题。实践中 Pre-LN 的训练稳定性明显优于 Post-LN。最后组装完整 ViTclass ViTForClassification(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, num_classes10, embed_dim192, depth6, num_heads3): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.blocks nn.Sequential(*[ TransformerBlock(embed_dim, num_heads) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) def forward(self, x): B x.shape[0] x self.patch_embed(x) cls self.cls_token.expand(B, -1, -1) x torch.cat([cls, x], dim1) x x self.pos_embed x self.blocks(x) x self.norm(x[:, 0]) return self.head(x)forward 里最后取出的是序列的第一个 token 位置也就是 class token而不是平均池化所有 patch。如果换成年平均池化也能用但 image 分类效果略差尤其是对于没有明显中心目标的图片。3.3 训练循环训练循环不算复杂但需要注意 loss 的计算和梯度回传ViT 在迭代上的行为与 CNN 有微妙差别——它更依赖 warmup所以训练循环要单独预留出调整学习率的逻辑。def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) correct (logits.argmax(dim1) labels).sum().item() return total_loss / len(loader.dataset), correct / len(loader.dataset) def evaluate(model, loader, criterion, device): model.eval() total_loss, correct 0, 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) logits model(images) loss criterion(logits, labels) total_loss loss.item() * images.size(0) correct (logits.argmax(dim1) labels).sum().item() return total_loss / len(loader.dataset), correct / len(loader.dataset)注意 evaluate 里用了torch.no_grad()这一步不可少它会关闭 autograd 的图构建。否则推理时每个前向都会保存中间激活内存占用随 batch 数线性上升。argmax 计算准确率时要注意维度位置logits 的形状是 (B, num_classes)argmax(dim1) 取出每个样本预测类别。优化器用 AdamW学习率先给一个小值后续章节再讨论怎么调model ViTForClassification().to(cuda) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) criterion nn.CrossEntropyLoss(label_smoothing0.1) for epoch in range(100): tr_loss, tr_acc train_one_epoch(model, train_loader, optimizer, criterion, cuda) te_loss, te_acc evaluate(model, test_loader, criterion, cuda) print(fepoch {epoch:3d} | train acc {tr_acc:.4f} | test acc {te_acc:.4f})固定 100 个 epoch 不动调参也能在 CIFAR-10 上达到 70% 以上的准确率但这个数字不是终点。label_smoothing 会让模型不再追求在每个样本上输出接近 one-hot 的极端概率减轻过拟合。如果发现训练准确率高但测试准确率停滞优先检查增强策略和 weight_decay而不是急着加深模型。4. 训练 ViT 的参数怎么设优化器选择、学习率策略与数据增强ViT 的训练配方相比 CNN 更敏感同一套超参数在不同数据集上的表现差异很大。官方训练配方涉及大批次和数据增强单卡环境下没有条件完全复刻需要针对显存做调整。这一章从优化器、学习率和增强三个方面展开给出参数表再讲调参依据。4.1 AdamW 和 weight decay 是 ViT 训练里的首选配置ViT 官方实现使用 AdamWweight decay 在 0.05 到 0.3 之间。相比之下传统 CNN 常用带动量的 SGD配合 0.0001 的 weight decay。为什么 ViT 更依赖 AdamWTrick 是 ViT 的归一化层LayerNorm对权重衰减敏感直接对 LayerNorm 的参数施加过大的 weight decay 会让训练不稳定所以需要精细控制。AdamW 把 weight decay 和梯度更新解耦实现上更干净。超参数CIFAR-10 / 单卡 128ImageNet / 多卡优化器AdamWAdamWlr1e-31e-3batch_size1284096warmup epochs510weight_decay0.050.3lr decayCosineCosinelabel_smoothing0.10.1dropout0.10.0表格里左右两栏的学习率都是 1e-3这个看似巧合的数值背后有一条线性缩放规则lr 大致随 batch size 线性增长。batch 128 时 lr 取 1e-3batch 翻倍到 256 时 lr 可以取 2e-3但 4096 这个量级不能再沿这条线推需要结合 warmup 调整。模型内部的 Dropout 设置要区分位置。PatchEmbed 后不加 DropoutEncoder 里的 Attention 和 MLP 都带 Dropout最后一层分类头前也有一层。在 CIFAR-10 上建议保留 dropout0.1因为数据量小、极易过拟合在 ImageNet 这种大规模数据上官方甚至不挂 dropout靠数据增强和正则化就够了。4.2 Cosine 学习率调度和 warmup 的实际作用import math def cosine_schedule_with_warmup(epoch, total_epochs, warmup_epochs5, lr_max1e-3, lr_min1e-5): if epoch warmup_epochs: return lr_max * (epoch 1) / warmup_epochs progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return lr_min 0.5 * (lr_max - lr_min) * (1 math.cos(math.pi * progress))warmup 阶段的学习率从很小开始线性升至 lr_max。原因是 ViT 的结构在训练初期非常容易发散patch embedding 和 class token 都是随机初始化的Attention 权重在早期梯度噪声很大直接施加一个大的学习率很容易让 loss 冲到 NAN。预热 5 个 epoch 是单卡 CIFAR-10 上的一个安全取值数据量大时按总步数 5% 的比例来定更标准。warmup 结束后接余弦衰减学习率先缓降再加速下降最终收敛到 lr_min。余弦调度的优势是全程变化平滑不像阶梯衰减那样带来尖峰式抖动。对比 step decay 和 cosine 在小数据集上的收敛曲线cosine 通常能多出 1~2 个点的准确率这一结论在 ViT 类模型上尤为明显。4.3 数据增强和 mixup 的适配原则ViT 在 CIFAR-10 上不加增强容易过拟合加了增强又可能因为增强太强而学不动。推荐的组合是RandomCrop HorizontalFlip 作为基础加上 RandAugment必要时再加 CutMix。transform_strong transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.RandAugment(num_ops2, magnitude9), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ])RandAugment 的 num_ops 表示每次随机选几个增强算子magnitude 表示增强强度。对 CIFAR-10 这类小图来说magnitude 取 79 比较合理太强会破坏图像本身的结构。CutMix 在软件实现上需要重写 loss 计算不能在 DataLoader 的 transform 里直接做因为它要破坏样本对的标签。4.4 训练不收敛时按这个顺序排查ViT 训练常见问题可以归成三类排查顺序建议是从下往上现象可能原因处理方式loss 一开始就 NANlr 过大 / fp16 overflow检查 lr换 bf16看 LayerNorm epsloss 不降位置编码维度不匹配检查 pos_embed 的第二维训练 acc 高、测试 acc 停滞数据增强不足增加 CutMix、RandAugment 强度测试时 OOM推理未关梯度检查是否漏了 torch.no_grad第一类最麻烦的是 fp16 训练下遇 NAN通常不是模型问题而是精度溢出。LayerNorm 在 fp16 下用 eps1e-6 容易导致数值不稳定建议改成 1e-5 或直接切换到 bf16二者在训练稳定性上差别很大。第二类原因常见于自己改 patch_size 后忘了同步 num_patches导致加位置编码时维度全部错位这类 bug 的报错信息通常很直接。第三类无关架构问题是训练数据太少下面一章专门谈。5. 小数据集上让 ViT 好好训练的 4 个实用技巧初始化、混合精度与蒸馏信号前面已经能够跑通并调参这一章压缩到最后四个最关键的生产级操作。在 CIFAR-10、森林图像分类或花卉分类这类千级别数据集上随机初始化的 ViT 几乎不可能打败同参数量的 CNN。常见的解法有两条路线用更小的模型或者引入预训练权重。第一个技巧是在模型初始化上动手脚。class token 的位置编码用零初始化会让训练早期梯度无法流入 Attention 的某一列改用标准差 0.02 的截断正态分布能更快进入正常收敛。patch embedding 的卷积层不要用默认的均匀初始化用 PyTorch 的 trunc_normal_ 会使 patch 之间有更合理的初始分布。第二个技巧是给 LayerNorm 的 eps 设大一点这在混合精度下尤其关键。坐标归一化层在 fp16 下过小的 eps 会导致除法结果溢出为 inf进而污染整个前向过程。ViT 官方实现使用nn.LayerNorm(dim, eps1e-6)但在小数据集和低精度训练中建议改成eps1e-5。配合 bf16 训练CIFAR-10 上的 loss 曲线会稳定很多。第三个技巧是用原规模的预训练骨架做特征迁移。这里不要求从零训练而是把 ViT 当作特征提取器冻结前若干层只训练分类头。实现时把每个 TransformerBlock 的 requires_grad 设成 False只保留最后一层可训练。这种方法在 dataset 只有几百张图时带来的收益常常比换训练时长更大。必要时可以只用全局平均池化替代 CLS token把提取出来的特征喂进线性分类器。第四个技巧是从 DeiT 里抄一个 distillation token。这个 token 的作用是在训练时额外引入一个监督信号比如让较小的 ViT 去模拟一个强大的 CNN 教师模型的输出。实现上只需要在序列末尾再拼一个 token训练 loss 里加入教师模型的 KL 散度项推理时仍然使用 class token 的输出。它的价值不在结构而在监督信息良好的教师能让学生快速收敛到稳定态。如果训练资源非常有限可以把 patch_size 从 16 改为 8 或 4。patch 越小序列越长计算量按平方增长但分类效果往往越好。动手前先用torch.profiler看一眼前向耗时序列长度为 196patch 16和 3136patch 4之间差 16 倍跑不跑得动在几秒内就能判断不值得为了追求极限效果把训练时间拖到不可接受。本文还有配套的精品资源点击获取
返回列表