
简介基于Vision TransformerViT实现CIFAR-10图像分类的训练与验证Python源码适合计算机视觉、人工智能方向的学生、教师及入门开发者参考学习。内容围绕ViT模型搭建、CIFAR-10数据预处理、训练与验证流程展开可直接运行用于图像分类实践也可作为毕设、课设或课程项目的初期演示基础。压缩包共2个文件包含一个Python主代码文件和一个项目结构说明txt文件整体体积仅2KB便于快速浏览与修改。目前已有525人学习下载代码经作者测试运行成功具备较好的可用性与完整性。解压后可通过结构说明快速了解项目文件组织便于在此基础上扩展实验或调整模型结构。1. 让ViT在小图上跑起来CIFAR10就是那块最好的试金石把 Vision TransformerViT用到 CIFAR10 上第一反应往往是“杀鸡用牛刀”——这个数据集每张图只有 32×32 像素而 ViT 的注意力机制天生擅长捕捉全局关系在小分辨率上既没有计算优势也容易过拟合。但正因为如此CIFAR10 才是验证 ViT 实现是否正确、训练技巧是否到位的“照妖镜”模型结构有没有写错、位置编码有没有加对、学习率策略是否合理都会在小数据集上暴露得干干净净。这篇文章要做的就是用 PyTorch 从零实现一个可训练的 ViT 分类模型跑通 CIFAR10 的完整训练和验证流程并给出踩坑记录。适合已经会用 PyTorch 做 CNN 分类、想转向 Transformer 架构的从业者以及需要一份能直接改用的最小可用 ViT 代码做 baseline 的人。2. ViT 结构拆解在动手写代码前先把三个关键组件想清楚2.1 Patch Embedding为什么 32×32 的图只能切 4×4 的 patchViT 的核心操作是把图像切成固定大小的 patch然后线性投影成 token 序列。CIFAR10 的图像是 32×32×3如果照搬原论文的 16×16 patch只能切出 2×24 个 token序列长度太短Transformer 的注意力机制根本施展不开而且每个 patch 尺寸过大位置信息几乎被抹平了。常见做法是把 patch size 设为 4×4这样能得到 8×864 个 token够用又不至于让序列太长拖慢训练。Patch Embedding 的实现不需要手动切图和拼接直接用卷积就能一步完成。nn.Conv2d(in_channels3, out_channelsembed_dim, kernel_sizepatch_size, stridepatch_size) 的效果等价于“切 patch 展平 线性投影”卷积核在图像上滑动步长等于 kernel_size每个输出位置恰好对应一个 patch 的线性映射结果输出的形状是 [B, embed_dim, H/patch, W/patch]再展平成 [B, embed_dim, num_patches] 并转置成 [B, num_patches, embed_dim] 即可。这里有个容易忽略的细节patch 的排列顺序必须保持从左到右、从上到下。卷积输出的空间维度天然满足这个顺序但如果你自己用 unfold 或 reshape 去切就得手动确认顺序否则位置编码和 token 的对应关系会错位模型训练出来精度会莫名其妙地差。实际工程中我一般直接写卷积简单可靠。2.2 Position Embedding 与 CLS Token两个必须同时处理好的细节ViT 的输入序列由三部分组成patch token 序列、一个可学习的 CLS token、以及加在两者之上的位置编码。CLS token 是放在序列最前面的一个可学习向量最终分类时只取它在 Transformer 输出中对应的那个向量接一个全连接层做分类。为什么需要它因为 Transformer 的输出是所有 token 的加权融合如果不额外加一个“用于分类”的 token就得对全部 patch token 做池化这样的全局平均会稀释掉少数关键 patch 的影响分类效果通常不如 CLS token 方案。位置编码这里有两个坑。第一个坑是在拼接 CLS token 之后位置编码的长度也要跟着加一。假设有 64 个 patch token拼上 CLS token 后序列长度是 65位置编码的 shape 必须是 (1, 65, embed_dim)写成 64 会直接报维度错误而新手经常写成 64 然后一头雾水。第二个坑是位置编码的初始化方式。用 PyTorch 默认的 nn.Parameter(torch.zeros(...)) 初始化训练前期 loss 会降得很慢因为所有位置初始完全相同注意力机制在一开始学不到位置差异。我一般用 nn.Parameter(torch.randn(...)) 配合一个小标准差初始化或者直接用 trunc_normal_ 初始化训练收敛明显快一些。2.3 Encoder 层手写多头注意力还是直接调 nn.MultiheadAttentionTransformer Encoder 层由多头自注意力、MLP、LayerNorm 和残差连接组成。这里有两种实现路线手写多头注意力逻辑或者直接调用 nn.MultiheadAttention。手写的优势在于可控性强你可以清楚看到 QKV 是怎么拆分的、mask 是怎么加的调试的时候能直接 print 出中间张量的形状。缺点是要注意的细节多容易在维度变换上出错——比如多头注意力的 head 维度拆分常见错误是把 [B, seq_len, embed_dim] 直接 reshape 成 [B, seq_len, num_heads, head_dim] 之后忘了转置成 [B, num_heads, seq_len, head_dim]导致后续矩阵乘法维度对不上。直接调 nn.MultiheadAttention 的优势是接口成熟、性能好PyTorch 内部的实现做了优化。但它也有坑这个模块默认的输入格式是 [seq_len, B, embed_dim]batch 在第二维和 Transformer 通用的 [B, seq_len, embed_dim] 不一样需要在传入前转置拿到输出后再转置回来。另外它默认会做 attention mask 和 key padding mask 的检查虽然 CIFAR10 分类不需要 mask但如果不小心传入了错误形状的 mask报错信息会绕得人头疼。下面的代码是我在生产环境中常用的 Encoder 层实现用的是手写多头注意力便于阅读和二次修改import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim, num_heads, dropout0.1): super().__init__() assert embed_dim % num_heads 0, embed_dim 必须能被 num_heads 整除 self.embed_dim embed_dim self.num_heads num_heads self.head_dim embed_dim // num_heads self.scale self.head_dim ** -0.5 # 将 QKV 投影合并成一个 Linear可以一次算完减少参数个数和计算开销 self.qkv nn.Linear(embed_dim, embed_dim * 3) self.proj nn.Linear(embed_dim, embed_dim) self.attn_drop nn.Dropout(dropout) def forward(self, x): B, N, C x.shape qkv self.qkv(x) # [B, N, 3*C] qkv qkv.reshape(B, N, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # [3, B, num_heads, N, head_dim] q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * self.scale attn F.softmax(attn, 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) return x逻辑说明QKV 投影合并成一个大 Linear 是 ViT 实现里最常见的写法这样 forward 里只需要一次矩阵乘法就能拿到完整的查询、键、值三个张量。reshape 和 permute 的步骤要仔细对照先把最后一维拆成 3Q/K/V、num_heads、head_dim 三组然后用 permute 把维度顺序调成 [3, B, num_heads, N, head_dim]这样 index 0/1/2 分别就是 q/k/v后续注意力计算不需要再改变布局效率更高。参数说明embed_dim 是整个模型的隐藏维度一般设为 192 或 256 就够 CIFAR10 使用num_heads 通常取 6 或 8这样每个 head 的 head_dim 在 24~32 之间太小会导致单个头的表达能力不足。dropout 在 CIFAR10 上建议设 0.1 左右数据集小dropout 太高会欠拟合。3. 数据准备与增强策略小图数据集上增强比模型结构更影响最终精度3.1 官方数据集下载与目录组织torchvision 接口和手动下载两种路径CIFAR10 数据集可以通过 torchvision 直接下载但国内网络环境下经常卡在下载阶段。如果你遇到 torchvision 下载超时的问题不要反复重试直接去数据集官网下载三个压缩包cifar-10-python.tar.gz 约 170MB手动放到 torchvision 默认查找的目录下即可。torchvision 会自动识别已存在的文件并跳过下载。目录组织方式如下项目根目录建 data/ 文件夹手动下载的压缩包直接放到 data/ 下torchvision 的数据集类会自动解压并整理。如果你的代码跑在服务器上强烈建议先把数据集下载好再上传服务器避免在训练节点上下载导致超时。数据加载的核心代码和训练/验证时的 transform 设计如下import torch from torchvision import datasets, transforms # 训练集增强 train_transform transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize( mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616] ), ]) # 验证集不做随机增强只做归一化 val_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize( mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616] ), ]) train_dataset datasets.CIFAR10( rootdata/, trainTrue, downloadTrue, transformtrain_transform ) val_dataset datasets.CIFAR10( rootdata/, trainTrue, downloadTrue, transformval_transform ) train_loader torch.utils.data.DataLoader( train_dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue ) val_loader torch.utils.data.DataLoader( val_dataset, batch_size256, shuffleFalse, num_workers4, pin_memoryTrue )逻辑说明训练集用 RandomCrop(32, padding4) 先把 32×32 的图像外扩到 40×40再随机裁剪回 32×32等效于给模型提供平移不变性的先验。RandomHorizontalFlip 对 CIFAR10 有效因为它的类别里没有“左右翻转后语义改变”的类别比如数字 6 和 9 翻转后会变但 CIFAR10 没有这类问题。参数说明mean 和 std 是 CIFAR10 数据集的通道统计量这三个值必须用官方计算好的不能用 ImageNet 的 mean/std 代替。如果用了 ImageNet 的归一化参数模型收敛会很慢且精度明显下降——原因是数据分布被错误地标准化了这不是玄学是数学上可以验证的偏差。num_workers 在本地建议设 2在服务器上设 4 或 8如果你的 CPU 核心数少设太高反而会因为进程切换开销拖慢数据加载速度。注意这里 val_dataset 使用的是 trainTrue 的 CIFAR10 数据集只是为了演示 transform 的拼接方式实际验证时必须用 trainFalse 加载测试集。训练集 50000 张、测试集 10000 张的划分是固定的不要把训练集混入验证。3.2 增强的边界CIFAR10 上哪些增强有效哪些适得其反在 CIFAR10 上用 ViT数据增强不是越猛越好。ViT 本身没有 CNN 的平移等变性先验所以对数据增强更敏感。我做过对比实验只用 RandomCrop 和 FlipViT 在测试集上大约能到 82%~85%加上色彩抖动ColorJitter后能提升 1~2 个点但如果加 CutMix 或 MixUp收敛速度会变慢需要配更长的训练轮数才能看到收益。具体来说色彩抖动要从轻微开始例如 brightness0.2, contrast0.2, saturation0.2 这样的一组参数。过强的色彩抖动会让模型学到的颜色特征失真因为 CIFAR10 里有不少类别高度依赖颜色信息比如青蛙和鸟形状相似但颜色差异大。如果抖动强度过大这些类别的区分度会被抹掉。AutoAugment 和 RandAugment 在 CIFAR10 上是有效的但要注意 ViT 对这种强增强的容忍度跟 CNN 不同。ViT 没有归纳偏置需要更多数据才能拟合因此强增强对它来说既是挑战也是机会。我一般用 timm 库提供的 RandAugment 实现配置 num_ops2, magnitude9能再提升 1 个点左右。但这些都是“最后一公里”的优化先把 baseline 跑通再调增强不要一上来就堆满增强否则翻车了都搞不清是模型问题还是数据问题。3.3 验证集上的评估流程哪些指标值得盯哪些指标会骗人CIFAR10 分类任务的验证指标以 Top-1 Accuracy 为主但只看准确率容易忽略细节。建议在验证集上额外计算每个类别的召回率因为 CIFAR10 类别不均衡度很低准确率基本能反映模型水平但如果某些类别召回率明显低于平均比如猫和狗经常互混说明模型学到的是形状特征而非关键语义特征这在后续迁移到真实数据时会吃亏。验证流程的代码要写成独立函数与训练循环解耦。验证阶段必须用 torch.no_grad() 包裹关闭梯度计算否则会额外占用显存并把 BatchNorm 的统计量搞乱——ViT 虽然不用 BatchNorm但 LayerNorm 也会受干扰。每次验证结束记录准确率和 loss用于判断模型是否过拟合。4. 训练循环与超参数配置CIFAR10 上把 ViT 训练稳定需要哪些关键设置4.1 优化器选择与学习率设置AdamW 是 ViT 训练的事实标准ViT 训练的事实标准是 AdamW 配合 cosine 学习率衰减。为什么不选 SGDViT 的 attention 层对学习率的敏感度远高于卷积层SGD 在 ViT 上训练极不稳定loss 曲线容易剧烈震荡。AdamW 的权重衰减是解耦的只对权重本身做衰减、不对梯度中的一阶动量做衰减这让正则化和优化互不干扰。学习率设置上ViT 一般用 lr1e-3 起步配 warmup。warmup 是必须的因为 Transformer 的 attention 层在训练初期梯度方差很大直接上大学习率会让 attention map 变成噪声。常见做法是前 10 个 epoch 线性从 1e-5 升到 1e-3之后 cosine 衰减到 1e-5。weight decay 在 CIFAR10 上建议设 5e-2这是 ViT 论文里的默认值但在小数据集上可以适当降到 1e-2防止正则化过强导致欠拟合。超参数的整体配置如下表所示这是我基于 CIFAR10 和 ViT-Tiny 的实际训练经验给出的基准值超参数建议值说明batch_size128显存 8GB 以下可降到 64base_lr1e-3AdamW 下的稳定起始点min_lr1e-5cosine 衰减的下限warmup_epochs10线性 warmup 的轮数weight_decay5e-2可下探到 1e-2total_epochs100CIFAR10 上建议至少 100 epochlabel_smoothing0.1减轻过拟合稳定验证精度4.2 完整训练循环一个能直接跑的 Python 脚本骨架下面的代码给出了从数据加载到训练再到验证的完整骨架可以直接复制保存成 train_vit_cifar10.py 运行import torch import torch.nn as nn from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR from tqdm import tqdm def train_one_epoch(model, loader, optimizer, criterion, device, epoch): model.train() total_loss, correct, total 0, 0, 0 pbar tqdm(loader, descfEpoch {epoch}) for images, labels in pbar: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() # 梯度裁剪ViT 训练早期容易出现梯度爆炸clip 能显著提高稳定性 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() * images.size(0) pred outputs.argmax(dim1) correct pred.eq(labels).sum().item() total images.size(0) pbar.set_postfix(lossloss.item(), acccorrect / total) return total_loss / total, correct / total def validate(model, loader, criterion, device): model.eval() total_loss, correct, total 0, 0, 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) pred outputs.argmax(dim1) correct pred.eq(labels).sum().item() total images.size(0) return total_loss / total, correct / total # 训练入口 device torch.device(cuda if torch.cuda.is_available() else cpu) model build_vit_tiny(num_classes10).to(device) # 结构见第 2 章 criterion nn.CrossEntropyLoss(label_smoothing0.1) optimizer AdamW(model.parameters(), lr1e-3, weight_decay5e-2) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) best_acc 0.0 for epoch in range(1, 101): train_loss, train_acc train_one_epoch( model, train_loader, optimizer, criterion, device, epoch ) val_loss, val_acc validate(model, val_loader, criterion, device) scheduler.step() print(fEpoch {epoch}: train_loss{train_loss:.4f}, ftrain_acc{train_acc:.4f}, val_acc{val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_vit_cifar10.pth) print(fBest val acc: {best_acc:.4f})逻辑说明train_one_epoch 和 validate 分离目的是让验证阶段不参与梯度计算和反向传播。梯度裁剪加在了 backward 之后、step 之前这个位置很关键因为裁剪针对的是梯度值本身必须在参数更新前完成。参数说明max_norm1.0 是梯度裁剪的阈值等于把所有参数的梯度 L2 范数限制在 1.0 以内超过的部分按比例缩放。ViT 在训练初期很容易出现个别层梯度范数超过 10 的情况不裁剪的话 loss 会突然变成 NaN。label_smoothing0.1 让模型不再追求对训练样本的绝对自信预测而是留一点概率给其他类别这在小数据集上能有效降低过拟合。4.3 混合精度与 Batch Size 的调参权衡ViT 的计算量比同参数量 CNN 大不少在 8GB 显存的卡上batch_size128 加上 256 hidden_dim 的 ViT-Tiny 勉强能跑。如果显存不够优先降 batch_size 而不是降模型维度因为 hidden_dim 降到 128 以下时注意力头的 head_dim 会太小影响精度。混合精度训练用 PyTorch 自带的 torch.autocast 和 GradScaler 即可。开启混合精度后显存占用大约能省 40%训练速度提升约 40%。但有一个坑LayerNorm 和最后的线性分类层在混合精度下容易出现精度损失所以 autocast 默认不会对这些层做降精度你不需要手动干预。如果用了自写的 LayerNorm记得检查是否支持 autocast否则数值稳定性会出问题。5. 避坑手册ViT 训练 CIFAR10 最常见的四个翻车现场5.1 位置编码与 patch 数量不匹配导致维度报错现象forward 第一次执行就报错错误信息是 The number of patches in the input sequence does not match the position embedding dimension。原因输入图像的高宽与 patch size 的整除关系被破坏。CIFAR10 是 32×32patch size 设为 4 没问题但如果你改了输入尺寸比如为了做测试把图 resize 成 224×224patch 数量变成 (224/4)²3136而位置编码形状还是 (1, 65, embed_dim)自然对不上。解决不要硬编码 patch 数量而是在 forward 里动态计算num_patches (H // patch_size) * (W // patch_size)然后切片或插值位置编码。写代码时用一个函数统一计算避免两个地方分别写死。5.2 训练 loss 降到 1.0 附近就不再下降验证精度停滞在 50%现象训练刚开始 loss 下降正常到 1.0 左右就卡住验证精度在 50%~55% 晃悠接近随机猜10 类随机猜测的准确率约 10%50% 说明模型学到了一些全局模式但学不到判别性特征。原因最常见的是 warmup 没做初始学习率太大导致 attention 层训练崩坏陷入了较差的局部最优。ViT 对学习率的敏感度远超 CNN1e-3 的初始学习率如果没有 warmup 垫底前几个 step 的梯度更新就会把注意力权重的分布搞坏。另一种可能是位置编码初始化全零模型没有位置信息patch 之间的顺序感是缺失的。解决加线性 warmup前 10 个 epoch 把学习率从 1e-5 线性升到 1e-3位置编码初始化换成 nn.init.trunc_normal_(pos_embed, std0.02)。改这两个地方loss 就能正常降到 0.5 以下。5.3 验证精度比训练精度高 5 个百分点以上现象训练精度 85%验证精度 90%看起来像是“模型在未知数据上表现更好”但冷静一下这不正常。原因绝大部分情况不是模型玄学而是数据泄漏。最典型的错误是把测试集 transform 写错——比如验证时也用了 RandomCrop或者验证数据里混入了训练集CIFAR10 的 torchvision 接口默认 trainTrue 加载的是训练集如果不小心把 train 参数写错验证集就是训练集的子集精度自然虚高。另有一种可能是 dropout 和 attention dropout 在训练时开启了、验证时忘了关闭模型在训练时被随机丢弃了一些信息表现被压制了。解决验证时严格只做 ToTensor 和 Normalize不做任何随机增强确认 val_dataset 的 trainFalse在 model.eval() 后确认模型里 dropout 被关闭。PyTorch 的 Dropout 层在 eval 模式下会自动关闭但如果你用了自定义的前向逻辑里手写了 dropout 函数需要手动判断。5.4 混合精度训练时 loss 变成 NaN现象开启 torch.autocast 后训练到第 3~5 个 epochloss 突然变成 NaN然后永久卡死。原因混合精度下梯度下溢。ViT 的最后一层分类头输出的 logits 在 FP16 下可能数值范围过大交叉熵损失计算时出现 inf。另外如果网络里有自定义的 exp、pow 运算FP16 下容易溢出。解决先看是不是 loss 计算部分的精度问题——把 criterion 移到 autocast 上下文之外让它在 FP32 下计算或者用 autocast 的 dtype 参数强制输出 FP32。如果是某个自定义激活函数的问题用 torch.nn.functional 的稳定版本替换。还有一个非常隐蔽的原因GradScaler 的 scale 因子初始值过大或过小用默认初始化即可不要手动乱调。6. 验证与进阶从准确率到“模型到底在关注什么”拿到一个训练好的 ViT 模型验证流程不能停在打印一个准确率数字。CIFAR10 是研究 ViT 行为的好数据集因为它小、跑得快、可视化成本低。我每次训练完会做两件事第一用测试集算混淆矩阵找出模型容易混淆的类别对第二挑几张测试图像把 attention map 可视化出来看模型在分类时到底在“看”图像的哪个区域。混淆矩阵的实现很简单在验证循环里收集所有预测结果和真实标签用 sklearn.metrics.confusion_matrix 生成。CIFAR10 上最常见的混淆对是猫和狗、鸟和鹿这两组的特征都是“四条腿带毛”或“天上有翅膀”ViT 在没有归纳偏置的情况下尤其容易混淆它们。如果你发现混淆对的准确率低于整体平均 10 个百分点以上就该考虑在数据增强里加一点点形状扰动比如随机擦除或者把模型尺寸加大一档。待办事项里优先级最高的是改造成你自己的数据集。替换 CIFAR10 的步骤很简单把数据集类换成你的 ImageFolder 格式目录重写 transform把分类头改成你自己的类别数。但要注意你的自然图像尺寸大概率不是 32×32那么 patch size、位置编码大小、warmup 和总训练轮数都需要重新调整。ViT 在小数据集上的迁移能力本身有限如果目标数据量低于 1 万张建议先用预训练 ViT 做微调而不是随机初始化从头训练。我个人的习惯是每个新项目都先跑一遍 CIFAR10 全流程确认模型结构和训练脚本没有低级错误再迁移到真实数据。这样能避免在真实数据上排错时连“到底是模型问题还是数据问题”都分不清。希望这个方案能帮你在 ViT 上少走一些弯路。本文还有配套的精品资源点击获取