Vision Transformer实战:CIFAR10图像分类从零到95%准确率

发布时间:2026/10/7 4:52:21
Vision Transformer实战:CIFAR10图像分类从零到95%准确率
简介基于Vision TransformerViT模型实现CIFAR-10图像分类的Python源码面向人工智能、计算机及相关专业的学生和研究人员可快速上手视觉Transformer的训练与验证流程。资源包含1个Python脚本和1个文本说明文件压缩包仅2KB脚本涵盖数据加载、模型构建、训练与评估等关键模块文本文件提供目录结构或运行要点说明便于使用者对照理解。目前已有526人浏览学习适合作为课程设计、毕业设计或入门Transformer项目的参考资料。代码经测试运行稳定可直接复现CIFAR-10分类实验亦可在其基础上扩展其他数据集或模型结构帮助读者熟悉ViT原理与PyTorch实现细节。整份资源轻量精简却覆盖了从数据预处理到验证评估的完整链路能有效节省环境调试时间。1. ViT 在 CIFAR10 上做分类小图用 Transformer 也能拿到 95% 准确率CIFAR10 这种 32×32 的小图过去默认是 CNN 的主场但 Vision Transformer 把图切成 patch 序列后同样能把验证准确率稳定推到 90% 以上配置得当接近 95%。这份源码就是一套完整的基于 ViT 实现 CIFAR10 分类数据集的训练和验证 Python 实现Vit.py 里是完整模型定义数据加载、训练循环、验证脚本和 checkpoint 保存都已串好装好 PyTorch 直接开跑。对正在做毕设、课设的计算机和人工智能相关专业学生或者想在 CIFAR10 上把 ViT 完整跑通、看清每个模块怎么配合的从业者这份代码比对照论文从零写省下大量时间。项目是测试成功后才打包的下载后先看 README.md改对数据路径就能复现。2. ViT 模型核心拆解patch embedding 怎么把 32×32 图像变成序列2.1 为什么选 ViT全局注意力比层层堆卷积更适合这种小图CIFAR10 一共 10 类物体分辨率只有 32×32一张图里物体的占比往往不小。CNN 处理这种小图时第一层卷积只能覆盖 3×3 或 5×5 的局部区域要得到全局关系得靠网络一层层叠高、慢慢扩大感受野。ViT 的思路完全不同先把图切成一堆 patch每个 patch 都和其余全部 patch 直接计算注意力权重等于第一层就在整张图范围内做信息交互。判断“船”和“卡车”这类差异集中在整体轮廓的类别时这种全局建模的收敛路径更直接。全局注意力也不是免费午餐。attention 本身对 patch 排列顺序不敏感必须靠位置编码把空间信息补回来同时 ViT 没有 CNN 那种内置的局部先验参数在小数据集上更容易过拟合。这也是 CIFAR10 上跑 ViT 必须把 patch_size、depth、增强策略和正则化系数都手工盯一遍的原因不是装个模型就能自动收敛。下面直接从 Vit.py 的核心代码看实现。2.2 Vit.py 核心代码PatchEmbedding、位置编码与分类头Vit.py 里最核心的是三个东西PatchEmbedding、位置编码和分类头。先看 PatchEmbedding常见做法是用一个带步长的卷积一次完成“切块线性映射”import torch import torch.nn as nn class PatchEmbedding(nn.Module): 把 32x32x3 的图切成 4x4 的 patch并线性映射到 embed_dim def __init__(self, img_size32, patch_size4, in_chans3, embed_dim192): super().__init__() self.num_patches (img_size // patch_size) ** 2 # 8x864 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # (B, 192, 8, 8) x x.flatten(2) # (B, 192, 64) x x.transpose(1, 2) # (B, 64, 192) return x一个二维卷积kernel_size 和 stride 都等于 patch_size效果等价于先把图切成 N 个 4×4 的块再对每个块做线性投影。输出 8×8×192 的特征flatten 成 64 个 token每个 token 是一段 192 维的向量。用卷积而不是手动切块是因为 GPU 对卷积的算子优化远好于循环切片训练速度能差出好几倍。参数上img_size32 是 CIFAR10 固定分辨率patch_size4 决定序列长度是 (32/4)^264这个数字越大 attention 矩阵越大embed_dim192 是每个 patch 映射到的向量维度改大会整体抬高模型参数量。接下来的 ViT 主体要把 class token 和位置编码加进去class ViT(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, num_classes10, embed_dim192, depth6, num_heads3, mlp_dim384, dropout0.1): super().__init__() self.patch_embed PatchEmbedding(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02) self.pos_embed nn.Parameter(torch.randn(1, num_patches 1, embed_dim) * 0.02) self.drop nn.Dropout(dropout) encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardmlp_dim, dropoutdropout, activationgelu, batch_firstTrue) self.encoder nn.TransformerEncoder(encoder_layer, num_layersdepth) 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) # (B, 64, 192) cls self.cls_token.expand(B, -1, -1) # (B, 1, 192) x torch.cat([cls, x], dim1) # (B, 65, 192) x self.drop(x self.pos_embed) x self.encoder(x) cls self.norm(x[:, 0]) # 取 cls token return self.head(cls)class token 是 ViT 作者从 BERT 借来的设计一个可学习的向量拼在 64 个 patch token 前面经过 encoder 后取它对应的输出做分类等价于把整张图的信息聚合到一个向量上位置编码的形状是 (1, 65, 192)多的 1 留给 class token。初始化用 randn×0.02 是常见做法让早期位置编码幅度和 patch 特征相当不会在 attention 里造成个别位置主导。encoder 用 PyTorch 自带的 TransformerEncoderLayer 堆 depth 层内部是 multi-head attention FFN LayerNormGELU 激活比 ReLU 在小模型上更平滑。位置编码为什么用加法而不是 concat因为位置信息是作为偏置叠加在 patch 特征上不会额外增加序列长度加法是 Transformer 里的标准做法concat 会导致维度翻倍且信息冗余。另外如果用的 PyTorch 版本较早不支持 batch_firstTrue可以去掉该参数在 forward 里把 x 转成 (seq_len, batch, embed) 格式其余逻辑不变。mlp_dim384 是 FFN 中间层宽度一般取 embed_dim 的 2 倍左右太小拟合能力不足太大会拖慢训练。这里有一个面试或答辩常被追问的设计问题为什么分类用 class token 而不是直接对 64 个 patch token 做 mean pooling原作者在 ViT 论文里做过对比两者在 ImageNet 上差距很小但 class token 在小数据集上更稳因为它通过训练自己学会“汇总全局信息”的方式而不是无脑平均。class token 的可学习性在小规模数据上天然规避了某个 patch 主导整张图信息的问题。3. CIFAR10 数据准备与训练主循环从加载数据到保存 best 模型3.1 用 PyTorch 加载 CIFAR10归一化、增强与 DataLoader 配置环境方面python 3.8 以上、torch/torchvision 装好即可数据集由 torchvision 自动下载离线环境就手工把 cifar-10-python.tar.gz 放进 root 目录。CIFAR10 官方把数据集分成 50000 张训练图和 10000 张测试图最省事的划分方式就是用测试集当验证集这也是这份源码的做法。import torchvision import torchvision.transforms as transforms transform_train 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]), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616]), ]) train_set torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) test_set torchvision.datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) train_loader torch.utils.data.DataLoader(train_set, batch_size128, shuffleTrue, num_workers2, pin_memoryTrue) val_loader torch.utils.data.DataLoader(test_set, batch_size128, shuffleFalse, num_workers2, pin_memoryTrue)RandomCrop(32, padding4) 先把图放大到 40×40 再随机裁回和 RandomHorizontalFlip 组合是最基础的空间增强对 ViT 这种没有内置平移不变性的模型尤其重要不增强的情况下小 ViT 在 CIFAR10 上大概率过拟合。Normalize 的 mean/std 是 CIFAR10 全数据集的统计值直接抄即可换成 ImageNet 那套统计值虽然也能训但收敛会差一些。DataLoader 里 shuffle 只在训练集打开num_workers2 在 Windows 上别拉到 8不然经常卡在 worker 初始化pin_memoryTrue 给 GPU 训练省一点拷贝时间CPU 训练设不设无所谓。提示第一次运行 downloadTrue 时需要联网下载数据公司内网环境建议先手动下载后放到 ./data 目录避免反复超时。3.2 训练主循环与验证主循环AMP、余弦退火和 checkpoint 保存数据准备好之后直接进训练主循环。下面这段逻辑里包含混合精度、学习率调度和验证函数是这份源码里的核心执行链路import torch import torch.nn as nn from torch.cuda.amp import GradScaler, autocast def evaluate(model, loader, device): model.eval() correct 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) preds outputs.argmax(dim1) correct (preds labels).sum().item() return correct / len(loader.dataset) device torch.device(cuda if torch.cuda.is_available() else cpu) model ViT(img_size32, patch_size4, embed_dim192, depth6, num_heads3, num_classes10).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay5e-2) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) scaler GradScaler() best_acc 0.0 for epoch in range(100): model.train() total_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() with autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss loss.item() * images.size(0) avg_loss total_loss / len(train_set) val_acc evaluate(model, val_loader, device) if val_acc best_acc: best_acc val_acc torch.save({epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_acc: best_acc}, best_model.pth) scheduler.step() print(fepoch {epoch1:3d} | train_loss {avg_loss:.4f} | val_acc {val_acc:.2%})几个步骤串起来看optimizer.zero_grad() 清零上一轮的梯度autocast 的作用是把 forward 和 loss 计算切成 fp16减少显存和计算量而反向传播由 GradScaler 的 scale(loss).backward() 接管这样能避免 fp16 小数值梯度直接下溢step 之前先 scaler.step(optimizer)最后 scaler.update()。如果不小心忘了 scaler.update()loss 会在某个 epoch 之后突然变差属于 AMP 常见的低级坑。验证函数单独拎出来是有原因的model.eval() 会把 Dropout 关掉、LayerNorm 统计状态切到推理模式torch.no_grad() 让模型不保存计算图。这两个少一个都不行少一个验证集准确率偏低或者显存越顶越高。ViT 里没有 BN但 Dropout 在 train/eval 的行为差异非常明显我见过有人因为只写 model.eval() 没包 no_grad()导致 16G 显存验证到一半 OOM。checkpoint 保存策略以 val_acc 为基准只要高于历史最高就覆盖保存 best_model.pth同时带上 optimizer_state_dict方便中断后从断点续训最后一个 epoch 再额外存一份 last_model.pth用来对比训练末期是否过拟合。不要只存最后一轮万一最后 10 轮过拟合了best 权重还在之前的位置。4. 超参数设置与验证指标CIFAR10 上跑 ViT 的几个关键数字4.1 推荐超参数表patch_size 和 depth 是关键开关这份源码在 CIFAR10 上验证过的配置我整理成参数表照着改基本不会翻车参数推荐值说明img_size32CIFAR10 固定分辨率patch_size4切成 8×864 个 patchembed_dim192比 256 省显存小数据集够用depth6增加层数收益有限显存开销大num_heads3192/364每个 head 64 维mlp_dim384和 embed_dim 保持 2 倍左右dropout0.1默认 0.1过拟合再加大batch_size1288GB 显存起步不够降到 64lr1e-3AdamW 配余弦退火用 1e-3 起weight_decay5e-2比 CNN 常用值大压过拟合epochs100小 ViT 在 CIFAR10 上 100 轮够看趋势schedulerCosineAnnealingLReta_min 设 1e-5收敛更稳先解释表格里最容易踩的两个参数。patch_size4 不是随手拍的32×32 的图切成 4×4 得到 64 个 patch序列长度和许多 NLP 短句差不多Transformer encoder 处理它不吃力改成 8 只剩 16 个 patch模型对空间细节的感知会明显钝化CIFAR10 准确率通常掉 23 个点。depth6 是这个小模型容量下的甜点区堆到 12 层在小数据集上并不线性涨点显存却是成倍往上跳。参数之间是联动的。embed_dim192、num_heads3 是配合关系num_heads 必须能整除 embed_dim否则 attention 分头时维度对不上会直接报错mlp_dim 取 embed_dim 二倍左右是为了让 FFN 有足够非线性变换空间。weight_decay5e-2 比 CNN 常用值大一个量级原因是 ViT 没有卷积归纳偏置attention 矩阵和偏置项都更需要权重衰减来压如果验证集曲线抖动先把 weight_decay 往上调而不是动 lr这是调参过程中最常被忽略的一条。4.2 验证函数与训练日志怎么判断模型真的训起来了训练日志建议每轮固定打印这几个字段方便对齐曲线epoch5/100 | train_loss0.8712 | val_loss0.9023 | val_acc71.2% | lr8.3e-4怎么判断train_loss 和 val_loss 同步下降val_acc 稳步爬升说明模型在学真东西train_loss 降、val_loss 不降甚至涨说明开始过拟合两个都不怎么动先确认是不是数据加载错了再考虑 lr 是不是太低。另外两个常见现象要说清楚ViT 在 CIFAR10 上第一轮 val_acc 在 40% 左右是正常的如果前 10 轮 val_acc 卡在 10%直接往数据归一化查别继续跑那是浪费电。还有一个比看数字更直观的技巧用训练集里的同一批固定图片每 10 轮做一次模型预测并保存结果图比只看 acc 更容易看出模型是在学边缘还是学背景色。这也是后面做 attention 可视化的前奏很多隐藏问题在图上是一眼就能瞄出来的。5. 常见问题与排查ViT 在 CIFAR10 上容易踩的 5 个坑下面这 5 个问题是我自己跑这份代码时真实遇到过的每一条都按现象、原因、解决三段给出来排查顺序按出现频率排越靠前越常见。5.1 验证准确率一直卡在 10% 附近现象loss 在 2.3 附近横盘val_acc 稳定在 9%11%跟随机猜的概率一模一样。原因九成是训练数据没做 Normalize。像素值 0255 进网络和 ViT 的 LayerNorm、初始化尺度对不上梯度方向从一开始就是乱的。另一种常见原因是验证集 transform 错写成了带 RandomCrop 的版本增强发生在推理阶段标签和内容错位。解决先检查数据管道训练集保留增强验证集只做 ToTensorNormalize。然后看模型的 logits 和 label 形状是否是 (B,10) 和 (B,)如果 CrossEntropyLoss 直接报错说明标签错位不报错却卡 10% 的基本只剩归一化问题。把 Normalize 补上后第一轮 val_acc 通常能到 30% 以上。5.2 loss 变 NaN学习率或 AMP 溢出的排查顺序现象前几轮 loss 正常下降跑到十几轮突然变 NaNval_acc 一落千丈。原因lr1e-3 在 ViT 上不算激进但叠加 AMP 时容易出现 fp16 梯度溢出位置编码初始化幅度过大也会放大早期梯度概率低一些。解决先把 lr 降到 5e-4 重跑还出现 NaN就在 backward 之后加一句 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)这行代码放在 scaler.step(optimizer) 之前。如果最终锁定是 AMP 的问题直接把 GradScaler 去掉CIFAR10 这种规模全精度训练完全跑得动没必要为了省显存抗 NaN。这个排查顺序是血泪经验一项一项排除比一次性换优化器快得多。5.3 GPU 显存不够batch_size 怎么降才不亏精度现象batch_size128 起训8G 显存的卡直接 CUDA out of memory。原因patch 序列长度为 64attention 中间变量在 depth6 的情况下也有不小的累积再加上 AMP 只省部分显存8G 卡跑 128 的 batch 很容易顶满。解决优先把 batch_size 降到 64还超就把 depth 从 12 砍到 6num_heads 从 6 砍到 3embed_dim 降到 192。降 batch_size 对收敛影响最小砍 depth 次之改 patch_size 影响最后的精度顺序不要反。再不行才考虑换大显存卡。ViT 对 batch_size 的敏感度比 CNN 高调小后最好同步回调 lr以 64 为基准lr 降到 5e-4 比较稳。5.4 训练集 99%、验证集 70%过拟合怎么压现象train_loss 一路冲到 0.1 以下train_acc 接近 100%val_acc 卡在 70% 附近还上下震荡。原因ViT 没有 CNN 的局部先验参数自由度大50k 张 CIFAR10 图片不够它把语义和噪声分开Dropout 和 weight_decay 不够时网络会把训练集细节直接背下来。解决默认 weight_decay5e-2、dropout0.1 起步已经过拟合就在 transform 里加 Cutout随机挖掉 16×16 的区域低成本高收益更有效的是 RandAugment。模型侧还可以加 Stochastic Depth也就是随训练轮数随机丢弃部分 encoder 层PyTorch 官方 TransformerEncoder 没直接实现可以用 vit-pytorch 这类库或者改源码里的 encoder 循环。优先走数据和正则化路线CIFAR10 数据是固定的不要指望加数据量。5.5 改输入尺寸后位置编码维度报错现象把 img_size 从 32 改成 64forward 直接抛 RuntimeError报错里出现 tensor a 的维度 17 vs 65 之类和信息相关但都是位置编码维度对不上。原因pos_embed 在init里按 num_patches1 写死成 65输入尺寸一改patch 数变多位置编码少了一大截拼接时对不上。解决把位置编码维度全推导出来别写死数字临时要跑实验可以用双线性插值把旧位置编码 resize 到新尺寸pos_embed_old model.pos_embed # (1, 65, 192) new_len num_patches 1 pos_embed_new torch.nn.functional.interpolate( pos_embed_old.permute(0, 2, 1), sizenew_len, modelinear ).permute(0, 2, 1) model.pos_embed torch.nn.Parameter(pos_embed_new)插值后的位置编码带一点位置信息失真能应急但正式实验还是直接重训。遇到这类报错先定位 num_patches (img_size // patch_size) ** 2再对照 pos_embed 的第二维5 分钟能查完。6. 进阶技巧把 attention map 抽出来看 ViT 到底在关注什么6.1 用 hook 拿注意力权重并叠加成热力图很多人跑完 ViT 只看 val_acc但这模型最大的特点是可解释把 attention map 画出来能直观看到分类依据。做法是挂一个 forward hook 在每层 self_attn 上PyTorch 的 nn.MultiheadAttention 返回 (attn_output, attn_weights)把 weights 缓存下来。PyTorch 2.x 里 TransformerEncoderLayer 有时会走 fast path 不返回权重hook 之前先强制设 need_weightsTrue。import numpy as np import matplotlib.pyplot as plt import cv2 attn_cache [] def hook_fn(module, input, output): # MultiheadAttention 的 output 是 (attn_output, attn_weights) attn_cache.append(output[1].detach().cpu()) for layer in model.encoder.layers: layer.self_attn.need_weights True layer.self_attn.register_forward_hook(hook_fn) model.eval() with torch.no_grad(): images, _ next(iter(val_loader)) _ model(images[:1].to(device)) # 取最后一层 attention在 head 维度平均 attn attn_cache[-1][0].mean(dim0) # (65, 65) cls_attn attn[0, 1:].reshape(8, 8) # cls 对应的 64 个 patch heatmap cv2.resize(cls_attn.numpy(), (32, 32), interpolationcv2.INTER_CUBIC) # 还原归一化叠加显示 img images[0].permute(1, 2, 0).numpy() img (img * np.array([0.2470, 0.2435, 0.2616]) np.array([0.4914, 0.4822, 0.4465])).clip(0, 1) plt.imshow(img) plt.imshow(heatmap, alpha0.5, cmapjet) plt.axis(off) plt.savefig(attn_map.png, bbox_inchestight)mean(dim0) 是把 3 个头的注意力平均成一个矩阵cls_attn 取的是 cls 这一行因为分类只用到 cls token 的注意力reshape(8,8) 对应 8×8 patch 网格resize 到 32×32 是还原原图坐标。热力图上越红的位置就是 cls token 重点关注的 patch。如果注意力集中到目标物体上说明模型在学形状如果集中在四角或背景要警惕网络拿背景色或纹理当捷径。这个可视化做完再回去调数据增强和超参比只盯着 val_acc 瞎猜要准得多。我最早在 CIFAR10 上训 ViT 时跑完 100 轮只看准确率后来发现某个类别的分类依据是背景的颜色自己居然一直没发现。从那以后我每次训练完都先抽一张 attention map 看一眼再谈调参这个习惯帮我避开了很多“看起来很好”的假象。希望帮到你。本文还有配套的精品资源点击获取