ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于ViT的CIFAR-10图像分类:训练与验证Python源码详解

基于ViT的CIFAR-10图像分类:训练与验证Python源码详解 简介基于Vit实现CIFAR10分类数据集的训练与验证Python源码包是一份可直接运行的深度学习实践项目面向计算机、人工智能、自动化等相关专业的学生、教师与从业者适合期末课程设计、课程大作业或毕业设计等应用场景。项目以Vision Transformer为核心将图像切分为序列化patches后通过自注意力机制捕捉全局特征完整覆盖了CIFAR10数据加载、预处理、模型搭建、训练调参、验证与性能评估等关键环节。压缩包内共12个文件以8个Python脚本为主包含vit.py、patch_embed.py、encoder_block.py等模型定义模块以及train_cifar10.py训练脚本另附README.md说明文档和训练效果可视化.png图片整体仅137KB目录结构清晰、便于按模块学习。目前已有227人学习了该项目代码均经过调试测试可稳定运行既能帮助初学者理解深度学习模型训练流程也可作为进阶研究者扩展优化、探索分类性能的基线框架。1. 用 ViT 做 CIFAR-10 分类为什么值得跑一遍这个源码如果你最近在关注视觉 TransformerViT这条技术路线多半见过它在 ImageNet 上刷榜的新闻。但真正把 ViT 源码跑起来很多人第一选择不是 ImageNet而是 CIFAR-10图片只有 32x32类别只有 10 个单卡就能训练迭代一轮只要几分钟。这个标题里的“基于 ViT 实现 CIFAR-10 分类数据集的训练和验证 python 源码”就是把 ViT 结构、训练循环、验证逻辑打包成一个最小可运行的项目。适合两类人第一类是刚看完 ViT 论文、想用代码验证结构理解的新手第二类是做过 CNN 分类、想对比 ViT 和 ResNet 在小数据集上表现的工程师。跑通它你能看到 Patch Embedding、Transformer Encoder、分类头这些概念在代码里到底长什么样也能知道为什么 ViT 在 CIFAR-10 上容易过拟合、该怎么调。下面我按自己落地这类项目的习惯把结构、数据、训练和避坑一条条拆开讲。2. ViT 模型结构拆解从 Patch Embedding 到分类头的代码实现2.1 输入图片怎么变成 TokenPatch Embedding 层ViT 和 CNN 最大的分水岭是输入处理方式。CNN 用卷积核在图片上滑动天然保留局部空间关系ViT 直接把图片切成固定大小的 Patch每个 Patch 拉平成向量当作 NLP 里的 Token 送进 Transformer。CIFAR-10 的图片是 32x32x3常见实现会选 patch_size4这样切出 8x864 个 Patch每个 Patch 的维度是 4x4x348。这 64 个 Token 加上一个分类用的 [CLS] Token一共 65 个 Token 进入编码器。代码里 Patch Embedding 通常用nn.Conv2d实现而不是真的切片再 reshape。原因很简单一个 kernel_sizestridepatch_size 的卷积输出通道设为 embed_dim每个输出位置的值就等价于对应 Patch 的线性投影。这样做不仅快而且梯度计算更直接。下面是最常见的一份实现import torch import torch.nn as nn class PatchEmbed(nn.Module): 将图像切块并投影到 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 # 64 # 等价于对每个 patch 做线性投影 self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: [B, 3, 32, 32] - [B, embed_dim, 8, 8] x self.proj(x) # 展平成 [B, embed_dim, num_patches] 再交换维度 x x.flatten(2).transpose(1, 2) # [B, 64, embed_dim] return x这里的embed_dim128是常见选择对应 ViT-Base 的 768 来说小了很多这是因为 CIFAR-10 数据量只有 6 万张训练图维度过大反而容易过拟合。num_patches必须提前算好因为后面要拿它初始化位置编码。如果你的输入图尺寸不是 32 的倍数这段代码会直接报错所以生产里一般会在前面加一个 Resize 层强制把输入缩放到img_size。2.2 Transformer Encoder 与分类头的核心参数得到 Token 序列后还要做两件事加一个可学习的 [CLS] Token再叠加位置编码。位置编码有两种路线ViT 论文用的是可学习的nn.Parameter后来很多实现换成了 sincos 固定编码理由是数据少时可学习位置编码容易过拟合。在 CIFAR-10 这种小数据集上我一般建议用可学习编码但维度不要太大后面避坑章会细说。Transformer Encoder 部分不需要自己从零写直接调nn.TransformerEncoder或者用 timm 的Block都行。不过既然标题是源码学习自己写一个更能看清参数。下面是核心的 Encoder Block 和分类头class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, dropoutdropout, batch_firstTrue) 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) ) def forward(self, x): # 先 norm 再 attention这是 pre-norm 结构 x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x class ViTForCIFAR10(nn.Module): def __init__(self, img_size32, patch_size4, embed_dim128, depth6, num_heads8, num_classes10, drop0.1): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, 3, embed_dim) self.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, self.num_patches 1, embed_dim)) self.pos_drop nn.Dropout(drop) self.blocks nn.Sequential(*[ TransformerBlock(embed_dim, num_heads, dropoutdrop) 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) # [B, 64, 128] cls self.cls_token.expand(B, -1, -1) # [B, 1, 128] x torch.cat([cls, x], dim1) # [B, 65, 128] x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) # 取 [CLS] token 对应的输出 x x[:, 0] x self.head(x) return x这里的关键参数是depth和num_heads。CIFAR-10 上depth6、num_heads8已经足够再加深会显著增加过拟合风险。mlp_ratio4是 ViT 论文的默认值即 MLP 隐藏层是 embed_dim 的 4 倍。另外注意nn.MultiheadAttention的batch_firstTrue必须加上否则输入输出维度排列就不是[B, seq_len, dim]新手在这里翻车很常见。3. 环境准备与 CIFAR-10 数据加载跑通训练前的最后一道坎3.1 Python 环境与依赖安装torch、timm、tensorboard跑这个源码之前先把 Python 环境弄干净。我建议用 Python 3.8 以上PyTorch 2.x 都行。依赖就三样torch、torchvision、timm。timm 不是必须的但如果你想用现成的 ViT 预训练权重做迁移学习它比手写代码方便得多。安装命令很简单但要注意 CUDA 版本匹配。如果你机器上没有 GPU也可以用 CPU 跑只是 CIFAR-10 一个 epoch 要几分钟训练 50 个 epoch 会让人失去耐心。# 创建虚拟环境 python -m venv vit_env source vit_env/bin/activate # Windows 下是 vit_env\Scripts\activate # 安装 PyTorch根据你的 CUDA 版本选择 index-url pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 安装 timm 和 tensorboard pip install timm tensorboard注意timm库的版本更新很快不同版本里timm.models.vision_transformer的参数名有差异。如果你直接按老博客的写法timm.models.VisionTransformer在新版里可能报错。我一般先跑python -c import timm; print(timm.__version__)确认版本再用。TensorBoard 不是必需品但对于观察损失曲线和准确率非常有用训练时顺手记录一下后面调参会轻松很多。3.2 DataLoader 与数据增强的配置细节CIFAR-10 数据集本身只有 32x32torchvision 可以直接下载。常见做法是把训练集和验证集分开训练集做随机水平翻转和随机裁剪验证集只做归一化。注意 CIFAR-10 默认的图片是 PIL 格式ToTensor会把像素值归一化到[0,1]然后再用均值和标准差做标准化。CIFAR-10 的全局均值是(0.4914, 0.4822, 0.4465)标准差是(0.2470, 0.2435, 0.2616)这套值来自官方统计不要自己随机改。from torchvision import datasets, transforms from torch.utils.data import DataLoader # 训练集增强随机裁剪 水平翻转 train_transform transforms.Compose([ transforms.RandomCrop(32, padding4), # 先 padding 再裁剪相当于随机平移 transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) # 验证集只需归一化 val_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_dataset datasets.CIFAR10(root./data, trainTrue, transformtrain_transform, downloadTrue) val_dataset datasets.CIFAR10(root./data, trainFalse, transformval_transform, downloadTrue) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size128, shuffleFalse, num_workers4, pin_memoryTrue)RandomCrop(32, padding4)这里 padding 参数是 4 不是 2因为默认 padding 模式是常数填充填充大小是裁剪边界的扩展量。CIFAR-10 图像小padding 太小增强效果不明显。pin_memoryTrue在 GPU 训练时能减少 CPU 到 GPU 的传输时间但如果你用 CPU 训练开了反而没意义。num_workers在我的 4 核机器上设为 4 刚好设太大反而会因为进程切换导致变慢。4. 训练与验证全流程从损失曲线到准确率指标4.1 训练循环优化器、学习率调度与正则化CIFAR-10 上的 ViT 训练和 CNN 有明显区别ViT 需要相对较小的学习率、更长的 warmup以及更强的正则化。这里说的正则化不只是 Dropout还包括权重衰减Weight Decay和随机深度。先看最基本的训练循环import torch import torch.nn as nn from torch.cuda.amp import autocast, GradScaler device torch.device(cuda if torch.cuda.is_available() else cpu) model ViTForCIFAR10(img_size32, patch_size4, embed_dim128, depth6, num_heads8, num_classes10).to(device) criterion nn.CrossEntropyLoss() # ViT 的默认优化器配置权重衰减全部加到非 norm/bias 参数上 optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) # warmup 5 个 epoch 余弦退火 scheduler torch.optim.lr_scheduler.LambdaLR( optimizer, lr_lambdalambda epoch: warmup_and_cosine(epoch, warmup_epochs5, total_epochs50) ) scaler GradScaler() # 自动混合精度 best_acc 0.0 for epoch in range(50): model.train() train_loss 0.0 correct 0 total 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() train_loss loss.item() * images.size(0) _, predicted outputs.max(1) total labels.size(0) correct predicted.eq(labels).sum().item() scheduler.step() print(fEpoch {epoch1:02d} Loss {train_loss/total:.4f} Acc {100.*correct/total:.2f}%)AdamW的weight_decay0.05是 ViT 训练的标准配置但要注意 PyTorch 的 AdamW 默认会对所有参数做权重衰减包括 LayerNorm 的 bias 和 scale 参数。严格来说应该把 bias 和 norm 参数排除在衰减之外否则训练不稳。常见做法是把参数分成两组传入优化器下面的避坑章会展开。学习率调度这里我写了一个自定义 lambda 函数。warmup 在前 5 个 epoch 内让学习率从 0 线性升到峰值之后用余弦函数衰减到接近 0。ViT 没有 warmup 很容易在训练初期就发散原因是 Transformer 的梯度方差比 CNN 大得多学习率稍微高一点就会让位置编码和 query/key 矩阵产生剧烈震荡。4.2 验证循环计算 Top-1 Accuracy 和混淆矩阵验证循环比训练循环简单但有几个细节要注意第一必须用model.eval()关闭 Dropout第二要在torch.no_grad()下推理第三验证集不需要梯度所以不能用autocast里的GradScaler。我习惯在验证时顺便统计混淆矩阵这样能看出模型到底把哪几类搞混。from sklearn.metrics import confusion_matrix import numpy as np def validate(model, loader, device): model.eval() criterion nn.CrossEntropyLoss() val_loss 0.0 all_preds [] all_labels [] with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() * images.size(0) preds outputs.argmax(dim1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) acc np.mean(np.array(all_preds) np.array(all_labels)) cm confusion_matrix(all_labels, all_preds) return val_loss / len(loader.dataset), acc, cm val_loss, val_acc, cm validate(model, val_loader, device) print(fVal Acc: {100.*val_acc:.2f}%) print(cm)这里val_loss我除以的是len(loader.dataset)不是len(loader)因为 loss.item() 已经乘了 batch size求平均损失应该按样本数归一。混淆矩阵在 CIFAR-10 上特别有用你会看到猫和狗、鸟和鹿经常互相混这是数据本身的语义相似性导致的不是模型 bug。如果你的验证精度卡在某个值不动先看混淆矩阵里哪两类混淆最严重再针对性做数据增强比盲目调学习率有效。5. 用 ViT 训 CIFAR-10 的避坑指南小数据集上最容易翻车的 5 个点5.1 过拟合为什么 epoch 还没过半训练精度 100% 而验证精度只有 60%现象训练第一个 epoch 损失就降得飞快第三个 epoch 训练准确率已经超过 90%但验证准确率一直停在 60% 左右。这是 ViT 在 CIFAR-10 上的典型过拟合。原因ViT 没有卷积的局部性先验全靠注意力机制从数据中学习空间结构需要的数据量远大于 CNN。CIFAR-10 只有 5 万张训练图对 ViT 来说太小了。解决第一把depth从 6 降到 4embed_dim从 128 降到 96模型参数变少过拟合会明显缓解。第二把dropout从 0.1 提高到 0.3并且打开DropPath随机深度让训练时随机丢弃一部分 Block 的输出。第三数据增强升级从RandomCrop RandomHorizontalFlip换成RandAugment它能同时调整对比度、饱和度、平移等。我用RandAugment(n2, m10)之后同样的模型验证精度从 68% 涨到 74%。最后权重衰减从0.05提到0.1效果显著。5.2 学习率与 warmupAdamW 的默认参数不是万能药现象用了lr3e-3训练前几个 epoch 损失不降反升然后 NaN。或者 warmup 写了但没生效损失曲线前 10 个 epoch 剧烈震荡。原因ViT 的初始查询向量和位置编码的梯度量级很大学习率太高直接导致梯度爆炸。另外很多开源代码的LambdaLR写法有问题warmup阶段的乘数不是从 0 开始而是从lr本身开始等于没做 warmup。解决峰值学习率用1e-3起步最多别超过5e-3。warmup 的 epoch 数设置为总 epoch 的 10% 到 20%比如总 50 epoch 就 warmup 5 个。另一个容易被忽视的点batch_size会直接影响最佳学习率。如果 batch 从 128 改成 256学习率应该按照平方根比例放大也就是乘sqrt(256/128)≈1.41否则大 batch 下的梯度更平滑同一学习率会显得偏小。我见过不少人用 batch 256 却保持lr1e-3结果收敛速度变慢。还有一点torch.optim.lr_scheduler.optimizer.param_groups[0][lr]才是实际生效的学习率不管scheduler.step()放在哪个位置建议每次打印确认。5.3 位置编码可学习的还是 sincos 的维度大小又该怎么选现象模型能跑通但验证精度比等价 CNN 模型低 5 个百分点以上。我排查半天发现位置编码的初始化方差太大。原因nn.Parameter(torch.zeros(1, 65, 128))如果改成torch.randn且没有乘以 0.02位置编码初始值会直接淹没 Patch 特征注意力机制一开始就把位置信息当成了主要信号。解决位置编码初始化为正态分布标准差取0.02是 ViT 论文中的推荐。也可以直接写成nn.Parameter(torch.randn(1, 65, 128) * 0.02)。对于 CIFAR-10 的 32x32 输入patch_size4 得到 64 个 patch和 Imagenet 上的 14x14196 个 patch 相比少得多。此时位置编码的学习压力也小如果追求极致稳定可以用 sincos 固定编码把pos_embed设为requires_gradFalse。我对比过在 CIFAR-10 上训练 50 epoch可学习编码比 sincos 高 1% 左右但前提是学习率足够低。如果学习率偏高sincos 反而更稳。5.4 自动混合精度训练Loss 变成 NaN 或直接不收敛现象torch.cuda.amp开启后前几个 step 正常到某一步 loss 变成 NaN然后一直回不来。原因ViT 的 attention 计算里有softmax在 FP16 下分母可能溢出尤其是当 logits 数值较大时。另一个原因是GradScaler没有被正确调用scaler.update()放在了optimizer.step()之前。解决检查代码里是不是忘了用with autocast():包住 forward。正确的顺序是optimizer.zero_grad()-with autocast(): loss criterion(model(images), labels)-scaler.scale(loss).backward()-scaler.step(optimizer)-scaler.update()。如果 loss 已经 NaN先把GradScaler的init_scale改小比如GradScaler(init_scale2**8)但这只是临时手段。从根本上说FP16 训练建议在nn.MultiheadAttention里把need_weightsFalse加上因为返回的 attention 权重矩阵在 FP16 下会额外占用显存并且容易溢出。如果你的 PyTorch 版本较新autocast会自动处理大部分问题但还是建议在训练脚本里加一个 NaN 检查if not torch.isfinite(loss).all(): optimizer.zero_grad(); continue保住前几个 batch 的进度。5.5 随机种子为什么每次训练结果差 3 个百分点现象在同一个机器、同一个参数下连续跑两次验证准确率差别达到 2% 到 3%。原因PyTorch 默认是不固定随机种子的数据加载的 shuffle、Dropout、初始化都会引入随机性。对 ViT 这种高方差模型小数据上稍微不同的初始化就可能收敛到不同的局部最优。解决在训练脚本开头统一设置种子包括 Python 的random.seed、NumPy 的np.random.seed、PyTorch 的torch.manual_seed如果在 GPU 上还要设置torch.cuda.manual_seed_all。同时设置torch.backends.cudnn.deterministic True和torch.backends.cudnn.benchmark False。注意num_workers0时DataLoader 的子进程会继承父进程的随机状态所以还要在worker_init_fn里为每个 worker 重新设置不同的种子。我用torch.initial_seed()加 worker id 做拼接保证可复现又不让每个 epoch 的数据顺序完全一样。这样调参时对比不同学习率才有意义否则你会误把随机波动当成方法改进。6. 把 80% 精度往上提迁移学习、CutMix 与 EMA 的实战技巧如果前面这些你都调过CIFAR-10 验证精度大概落在 75% 到 80%。想再往上走单靠训练自己的小 ViT 很难因为它们没有预训练权重相当于从零学习空间特征。常见做法是加载在 ImageNet-21k 上预训练好的 ViT-Tiny 或 ViT-Small 权重然后只微调最后几层和分类头。这里有个关键点ImageNet 预训练的 patch_size 是 16位置编码对应 14x14 的 patch 网格CIFAR-10 图片是 32x32如果直接 resize 成 224 再切成 16p 的 patch等于把整个图放大 7 倍空间细节全没了。我一般会把预训练模型的位置编码用插值重采样到 8x8 的网格再把 patch_embed 的卷积核从 16 改成 4 并重新初始化这样输入 32x32 图片才匹配。这个改动在 timm 里就是timm.models.create_model(vit_tiny_patch16_224, pretrainedTrue, img_size32, patch_size4)但旧版本不支持这种运行时改法建议直接读源码改一下。除了迁移学习训练技巧里最能涨点的是 CutMix。CutMix 会把一张图的随机区域剪下来贴到另一张图上并且标签也按照面积比例混合。它比 Mixup 更适合 ViT因为 ViT 的注意力分布集中在 patch 上CutMix 能强制模型关注全局而非某一个区域。实现 CutMix 不需要额外库从 torch 官方仓库拷贝一段就行。配合 EMA指数移动平均让模型参数跟随训练过程中多次平均值推理时用 EMA 权重通常能再涨 1 到 2 个点。我在 CIFAR-10 上把这些做完从零训练的 ViT 精度从 76% 提到了 84%。这套流程跑通之后你要做的第一件事是把 MobileNetV2 或 ResNet-18 放在同样的数据增强和优化器配置下对比你会发现 CNN 在小数据上仍然更稳ViT 的价值在更大数据和更强算力下才更明显。这也是为什么很多新项目宁可用 Swin Transformer 这种带层级设计的变体也不直接上纯 ViT。希望这些经验能帮你少踩几个坑把源码跑得又快又好。本文还有配套的精品资源点击获取
返回列表