
简介这是一份面向深度学习初学者的Vision TransformerViT实战入门资源聚焦图像分类任务特别适合希望快速掌握Transformer在CV领域落地应用的开发者与学生。资源基于PyTorch实现以植物幼苗数据集12类为案例完整覆盖ViT模型构建、自定义数据集生成、Cutout与Mixup双数据增强策略、训练/验证流程、余弦退火学习率调度及两种预测写法等核心环节代码简洁无冗余便于逐行理解与复现。压缩包共2418个文件主体为2406张PNG格式植物幼苗图像用于训练与验证辅以7个功能清晰的Python脚本含模型定义、训练主逻辑、增强实现等及5个编译缓存文件整体大小930.96MB。已有3051人学习下载内容结构直击ViT工程实践关键节点提供可即用的数据组织方式、可调试的训练模板与可迁移的增强配置是少有的兼顾原理性与操作性的轻量级ViT教学资源。1. 项目概述为什么Vision Transformer值得你花时间如果你正在计算机视觉领域摸索或者对深度学习模型感兴趣那么Vision Transformer简称ViT这个名字你一定不陌生。它不像卷积神经网络那样需要你从卷积核、池化层这些概念开始慢慢搭建认知ViT直接把自然语言处理领域的“明星”Transformer架构搬到了图像上用一种近乎“暴力美学”的方式处理视觉任务。我第一次接触ViT时感觉就像有人告诉我“嘿别管那些复杂的局部特征提取了我们把图片切成小块然后像处理句子里的单词一样处理它们就行”。这种思路在当时看来非常大胆甚至有点反直觉但结果却出奇地好。这个项目或者说这篇总结就是为你准备的。无论你是刚学完CNN想看看新东西的学生还是工作中需要快速评估一个新模型潜力的工程师亦或是单纯对前沿技术好奇的爱好者这篇“非常简单的ViT入门教程”都试图做到一件事让你在最短的时间内亲手跑通一个ViT模型并理解它最核心的运作原理。我们不追求大而全的论文复现而是聚焦于“实战”和“入门”。我会带你从零开始用PyTorch搭建一个最基础的ViT模型在经典的CIFAR-10数据集上完成图像分类任务。过程中我会穿插解释每一个关键步骤背后的设计逻辑以及我在多次复现中踩过的坑和总结的技巧。相信我看完并跟着做一遍你会对ViT有一个扎实且直观的理解这远比读十篇抽象的论文综述来得有效。2. ViT核心思想拆解图像如何变成“句子”在深入代码之前我们必须先搞懂ViT最根本的思想。传统的CNN通过卷积核在图像上滑动来提取局部特征这种归纳偏置即先验知识图像中相邻像素关联性强是其成功的关键但也可能限制了模型学习更全局关系的能力。Transformer在NLP领域的巨大成功证明了其自注意力机制在捕捉长距离依赖关系上的强大能力。ViT的核心创新就在于思考我们能否抛弃卷积直接用Transformer来处理图像2.1 图像分块嵌入从像素到“视觉单词”Transformer的输入是一串序列比如单词的嵌入向量。图像是二维的网格第一步就是要把它“序列化”。ViT的做法非常直接分块将一张输入图像例如224x224像素3个通道分割成固定大小的、不重叠的块。论文中常用的是16x16像素的块。那么一张224x224的图像就会被分成 (224/16) * (224/16) 14 * 14 196个块。展平与线性投影每个块16x16x3768个像素值被展平成一个长度为768的向量。然后这个向量通过一个可训练的线性投影层一个全连接层映射到一个模型设定的隐藏维度D例如768。这个步骤就相当于NLP里把单词通过嵌入层变成词向量的过程。至此我们得到了196个长度为D的向量它们就是图像的“视觉单词”。注意这里的分块大小和投影维度是超参数。更小的块如8x8会产生更长的序列模型更精细但计算量剧增更大的块则相反。D的维度决定了模型表征能力的基础宽度。2.2 位置编码与可学习的分类令牌Transformer本身是置换不变的打乱输入序列顺序输出不变但图像中块的位置信息至关重要。因此我们需要为每个块向量添加位置编码。ViT使用标准的可学习的一维位置编码即模型自己学习一组与块位置对应的向量然后加到对应的块嵌入向量上。此外为了完成分类任务ViT借鉴了BERT的做法在序列的开头添加了一个额外的、可学习的向量称为分类令牌。这个CLS令牌会与所有图像块信息进行交互并在经过Transformer编码器后其对应的输出向量被用于最终的分类预测。你可以把它理解为一个特殊的“侦察兵”它跑遍整个图像“战场”后回来汇报整体战况图像类别。2.3 Transformer编码器自注意力的舞台处理完的序列CLS令牌 图像块令牌 位置编码被送入标准的Transformer编码器堆栈。每个编码器层主要包含两个核心子层多头自注意力层这是精髓。对于序列中的每一个令牌包括CLSMSA层允许它“关注”序列中的所有其他令牌包括它自己。通过计算注意力分数模型可以动态地决定在整合信息时应该更“重视”哪些块。例如要判断一张图片是不是“狗”模型可能会让CLS令牌更多地关注包含狗头、狗尾巴的块而忽略背景中的草地。前馈神经网络层一个简单的多层感知机通常包含两个线性变换和一个激活函数如GELU用于对每个令牌进行独立的、非线性的特征变换。每个子层后面都跟着层归一化和残差连接这是稳定深层模型训练的关键技术。2.4 归纳偏置的缺失与数据需求ViT最引人讨论的一点是它几乎摒弃了CNN固有的归纳偏置。它没有利用图像的2D结构、局部性、平移不变性等先验知识所有的空间关系都需要模型从数据中从头学习。这既是其强大之处理论上可以学习到任何形式的关系也是其“弱点”——它需要海量的数据才能充分训练。原论文中ViT在大型数据集如JFT-300M3亿张图像上预训练后迁移到下游任务如ImageNet上才能取得超越CNN的效果。如果只在ImageNet130万张图上从头训练其表现可能不如同等规模的CNN。这一点对于入门者至关重要它解释了为什么我们通常不会在小型数据集上从头训练一个大ViT而是采用预训练-微调的策略。3. 从零搭建一个简易ViT代码逐行解析理论说得再多不如动手写一行代码。下面我将用PyTorch一步步实现一个简化版的ViT用于CIFAR-1032x32小图像分类。我们会适当调整参数以适应小尺寸图像。3.1 环境准备与数据加载首先确保你的环境安装了PyTorch和Torchvision。我们使用CIFAR-10因为它体积小训练快非常适合教学和实验。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms import numpy as np # 定义数据预处理和增强 train_transform transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.RandomCrop(32, padding4), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), # CIFAR-10的均值和标准差 ]) test_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) # 加载数据集 train_dataset datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtrain_transform) test_dataset datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtest_transform) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size128, shuffleFalse, num_workers2)3.2 核心模块实现Patch Embedding这是将图像转换为令牌序列的模块。对于CIFAR-10的32x32图像我们使用4x4的块这样会得到(32/4)^2 64个块。class PatchEmbedding(nn.Module): 将图像分割成块并嵌入。 输入: (B, C, H, W) 输出: (B, num_patches, embed_dim) def __init__(self, img_size32, patch_size4, in_channels3, embed_dim128): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # 使用一个卷积层来实现“分块线性投影”效率更高 # 卷积核大小步长patch_size输出通道数embed_dim self.projection nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): B, C, H, W x.shape # 投影: (B, C, H, W) - (B, embed_dim, H/patch, W/patch) x self.projection(x) # 展平空间维度: (B, embed_dim, H/patch, W/patch) - (B, embed_dim, num_patches) x x.flatten(2) # 调整维度顺序: (B, embed_dim, num_patches) - (B, num_patches, embed_dim) x x.transpose(1, 2) return x为什么用卷积层而不是全连接层虽然论文描述是“展平后线性投影”但在实现上一个与块大小相同的卷积核、且步长等于块大小的卷积层其数学操作是完全等价的并且计算效率更高更符合图像处理的习惯。3.3 核心模块实现Transformer编码器层我们先实现多头自注意力MHA和前馈网络FFN然后组合成编码器层。class MultiHeadSelfAttention(nn.Module): 简化版多头自注意力未包含缩放和掩码因ViT是完整序列 def __init__(self, embed_dim128, num_heads8, dropout0.1): super().__init__() assert embed_dim % num_heads 0, embed_dim必须能被num_heads整除 self.num_heads num_heads self.head_dim embed_dim // num_heads self.scale self.head_dim ** -0.5 # 缩放因子稳定训练 self.qkv nn.Linear(embed_dim, embed_dim * 3) # 同时计算Q, K, V self.attn_dropout nn.Dropout(dropout) self.proj nn.Linear(embed_dim, embed_dim) self.proj_dropout nn.Dropout(dropout) def forward(self, x): B, N, C x.shape # N是序列长度块数1 # 计算Q, K, V: (B, N, C) - (B, N, 3C) - 3个(B, N, C) qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] # 每个都是(B, num_heads, N, head_dim) # 计算注意力分数: (B, num_heads, N, head_dim) (B, num_heads, head_dim, N) - (B, num_heads, N, N) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.attn_dropout(attn) # 应用注意力到V: (B, num_heads, N, N) (B, num_heads, N, head_dim) - (B, num_heads, N, head_dim) x (attn v).transpose(1, 2).reshape(B, N, C) # 合并多头 # 输出投影 x self.proj(x) x self.proj_dropout(x) return x class FeedForward(nn.Module): 简单的前馈网络两层线性层GELU激活 def __init__(self, embed_dim128, hidden_dim512, dropout0.1): super().__init__() self.net 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): return self.net(x) class TransformerEncoderLayer(nn.Module): 一个完整的Transformer编码器层MSA - Add Norm - FFN - Add Norm def __init__(self, embed_dim128, num_heads8, ff_hidden_dim512, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn MultiHeadSelfAttention(embed_dim, num_heads, dropout) self.dropout1 nn.Dropout(dropout) self.norm2 nn.LayerNorm(embed_dim) self.ffn FeedForward(embed_dim, ff_hidden_dim, dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x): # 残差连接1 x x self.dropout1(self.attn(self.norm1(x))) # 残差连接2 x x self.dropout2(self.ffn(self.norm2(x))) return x实操心得在实现自注意力时使用单个线性层同时计算Q、K、Vself.qkv然后通过reshape和permute进行拆分是常见且高效的技巧。注意LayerNorm是放在残差块内部的Pre-Norm这与原始Transformer的Post-Norm不同Pre-Norm通常能使深层模型训练更稳定。3.4 组装完整的ViT模型现在我们将所有部件组装起来并添加CLS令牌和位置编码。class SimpleViT(nn.Module): def __init__(self, img_size32, patch_size4, in_channels3, num_classes10, embed_dim128, depth6, num_heads8, ff_hidden_dim512, dropout0.1): super().__init__() self.patch_embed PatchEmbedding(img_size, patch_size, in_channels, embed_dim) num_patches self.patch_embed.num_patches # 可学习的CLS令牌和位置编码 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) # 1 for cls_token self.pos_dropout nn.Dropout(dropout) # Transformer编码器堆栈 self.encoder_layers nn.ModuleList([ TransformerEncoderLayer(embed_dim, num_heads, ff_hidden_dim, dropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) # 分类头 self.head nn.Linear(embed_dim, num_classes) # 初始化参数 nn.init.trunc_normal_(self.cls_token, std0.02) nn.init.trunc_normal_(self.pos_embed, 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.constant_(m.bias, 0) elif isinstance(m, nn.LayerNorm): nn.init.constant_(m.bias, 0) nn.init.constant_(m.weight, 1.0) def forward(self, x): B x.shape[0] # 1. 生成图像块嵌入 x self.patch_embed(x) # (B, num_patches, embed_dim) # 2. 添加CLS令牌 cls_tokens self.cls_token.expand(B, -1, -1) # (B, 1, embed_dim) x torch.cat((cls_tokens, x), dim1) # (B, num_patches1, embed_dim) # 3. 添加位置编码 x x self.pos_embed x self.pos_dropout(x) # 4. 通过Transformer编码器 for layer in self.encoder_layers: x layer(x) # 5. 取CLS令牌的输出用于分类 x self.norm(x) cls_output x[:, 0] # 取第一个令牌CLS的输出 # 6. 分类头 logits self.head(cls_output) return logits关键点解析cls_token和pos_embed都是可学习的参数随模型一起训练。在forward中cls_token需要根据批次大小B进行扩展然后与图像块序列拼接。位置编码直接加到令牌序列上。原论文使用的是固定的一维正弦编码但可学习的位置编码在实践中同样有效且更简单。最终只取序列第一个位置即CLS令牌的输出向量送入分类头。4. 模型训练、调优与问题排查模型搭好了接下来就是训练它。对于ViT这类数据饥渴型模型在CIFAR-10这样的小数据集上从头训练需要一些技巧。4.1 训练脚本与超参数设置device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleViT( img_size32, patch_size4, num_classes10, embed_dim128, # 隐藏层维度较小以适应小数据集 depth6, # Transformer层数 num_heads8, ff_hidden_dim512, dropout0.1 ).to(device) criterion nn.CrossEntropyLoss() # 使用AdamW优化器它对Transformer类模型效果更好并带有权重衰减 optimizer optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) # 使用余弦退火学习率调度有助于稳定训练并提高最终精度 scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50, eta_min1e-6) num_epochs 50 for epoch in range(num_epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() # 梯度裁剪防止梯度爆炸对Transformer训练很重要 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() running_loss loss.item() * images.size(0) scheduler.step() epoch_loss running_loss / len(train_loader.dataset) print(fEpoch [{epoch1}/{num_epochs}], Loss: {epoch_loss:.4f}, LR: {scheduler.get_last_lr()[0]:.6f}) # 简单验证 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() print(fTest Accuracy: {100 * correct / total:.2f}%)超参数选择心得学习率3e-4是训练Transformer的一个常用起点。对于小数据集可以尝试更小的值如1e-4。权重衰减AdamW配合权重衰减这里设为0.05是防止过拟合的关键尤其对于参数量较大的模型。优化器绝对不要使用普通的SGD或Adam一定要用AdamW。AdamW正确地将权重衰减与梯度更新解耦对于ViT的训练至关重要。学习率调度余弦退火在视觉Transformer中几乎是标配它能平滑地将学习率从初始值降到接近零有助于模型收敛到更好的局部最优点。梯度裁剪虽然不像在RNN中那么必要但加上梯度裁剪clip_grad_norm_是一个好习惯能增加训练稳定性。4.2 常见问题与排查技巧实录在实际复现ViT时你几乎一定会遇到下面这些问题。这里是我的排查笔记问题1损失不下降准确率随机约10%等于瞎猜现象训练了几个epoch损失值在高位震荡测试准确率始终在10%左右CIFAR-10有10类。可能原因与排查数据流错误首先检查数据预处理和加载。打印一个批次的图像和标签看形状和范围是否正确图像是否归一化到[0,1]或[-1,1]标签是否在0-9之间。模型前向传播错误在训练循环开始前用一组随机数据torch.randn(B, C, H, W)进行一次前向传播检查输出logits的形状是否为(B, num_classes)并且没有出现NaN或Inf。初始化问题ViT对初始化敏感。检查你是否正确初始化了cls_token和pos_embed使用trunc_normal_。确保所有线性层和LayerNorm层都按照_init_weights方法正确初始化。学习率过高/过低尝试一个更极端的学习率如1e-5或1e-3跑1-2个epoch看损失是否有任何变化。如果都没变化可能不是学习率的问题。我的踩坑记录我曾忘记将cls_token和pos_embed注册为nn.Parameter导致它们没有被优化器更新模型永远学不到有效特征。务必用print(model.named_parameters())检查所有你认为可学的参数是否在列表中。问题2训练后期过拟合严重现象训练损失持续下降但验证准确率在达到一个峰值后开始下降。对策增强正则化增加dropout率尝试0.2或0.3。增加weight_decay尝试0.1。对于小数据集这非常必要。使用更强的数据增强除了水平翻转和随机裁剪可以考虑加入AutoAugment、RandAugment或MixUp/CutMix等更现代的数据增强策略。这对ViT在小数据集上的表现提升显著。减少模型容量降低embed_dim如从128降到64、减少depth层数或num_heads。我们的简易ViT参数不多但在CIFAR-10上embed_dim128, depth6已经算“大”模型了。早停监控验证集精度当连续多个epoch不再提升时停止训练。问题3训练速度慢GPU内存占用高现象每个epoch耗时很长或者批次大小batch size无法设大。优化技巧减小序列长度这是最大的瓶颈。序列长度是num_patches 1。对于224x224图像16x16分块产生197的长度。自注意力的计算复杂度与序列长度的平方成正比。在入门阶段务必使用小图像如32x32、64x64和小分块如4x4、8x8来降低序列长度。使用混合精度训练PyTorch的torch.cuda.amp可以显著减少GPU内存占用并加速训练。梯度累积如果GPU内存不足以支撑大的batch size可以使用梯度累积。例如设置batch_size32但每4个批次才更新一次梯度累积步数4这等效于batch size128的效果但内存占用仅为1/4。问题4注意力图可视化一片模糊或没有意义现象想可视化自注意力权重来看看模型关注哪里但得到的图看不出任何模式。排查检查取出注意力的位置确保你是在模型训练收敛后取出某一层、某一头的注意力权重矩阵。注意力矩阵的形状应为(B, num_heads, N, N)其中N是序列长度。可视化CLS令牌对其他块的注意力通常分类任务中最有意义的是CLS令牌序列索引0对所有图像块令牌的注意力。将attn[0, :, 0, 1:]取第一个样本所有头CLS令牌对图像块的注意力进行平均或选择某个头进行可视化。模型可能没训练好如果模型本身性能就很差注意力图自然没有意义。先确保模型在测试集上有不错的准确率例如在CIFAR-10上达到80%。5. 超越入门下一步可以做什么当你成功运行了这个简易ViT并理解了其运作方式后你已经掌握了ViT最核心的骨架。要走向更实际、更强大的应用以下几个方向值得深入1. 使用预训练模型进行迁移学习这是ViT最实用的打开方式。Hugging Face的transformers库或timm库提供了在ImageNet-21k或JFT上预训练好的各种ViT模型如vit-base-patch16-224。你可以轻松加载它们只替换最后的分类头然后在你的小数据集上进行微调。这通常只需要很少的epoch和计算资源就能获得远超从头训练的效果。# 使用timm库示例 import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10) # 冻结除分类头外的所有参数可选 for param in model.parameters(): param.requires_grad False model.head nn.Linear(model.head.in_features, 10) # 替换分类头2. 探索ViT的变体与改进原始的ViT有很多已知的改进方向Swin Transformer引入了层次化构建和滑动窗口注意力恢复了卷积的局部性和层次性先验在多个视觉任务上实现了SOTA且计算复杂度线性增长。DeiT通过引入一个蒸馏令牌让ViT能够从CNN老师那里学习从而在不依赖海量数据的情况下仅用ImageNet就能训练出优秀的ViT。MobileViT致力于将ViT部署到移动设备在精度和效率之间取得了很好的平衡。3. 将其应用到其他视觉任务ViT不仅是分类模型其编码器可以作为强大的视觉特征提取器。目标检测将ViT作为DETR等检测模型的主干网络。语义分割将ViT输出的块特征重新排列成2D特征图再接上分割头如FPN、UPerNet。图像生成ViT也可以作为扩散模型或GAN中的核心组件。我个人最实际的建议是不要满足于在CIFAR-10上跑通。找一个你感兴趣的小型自定义图像数据集哪怕是爬虫爬的几百张图片尝试用预训练的ViT进行微调解决一个实际的分类问题。这个过程中遇到的图片尺寸调整、数据不平衡、过拟合等问题以及解决它们的过程会让你对ViT乃至深度学习工程有质的理解。从“跑通教程”到“解决自己的问题”这一步跨越才是学习的真正开始。本文还有配套的精品资源点击获取