资讯动态

Python+ViT实战:CIFAR-10图像分类从70%到96%的调优指南

发布时间:2026/10/5 4:09:17 来源:尧图企业网站定制
简介这份资源面向深度学习初学者与课程实践者提供一套基于Vision Transformer完成CAFIR10图像分类的完整大作业方案帮助读者理解如何用Python将图像切分为patch并借助自注意力机制实现全局特征建模适合作为课程设计、期末项目或Transformer入门练手素材。压缩包共21个文件约11.25MB包含7个ipynb实验笔记、3个py源码、3个docx文档、3个pptx汇报材料以及csv数据与txt说明覆盖从数据读取、模型搭建到训练评估与结果展示的完整链路。资源还涉及手写数字识别、机器翻译、LSTM自动写诗等同类作业模块便于横向对比不同任务的实现思路。目前已有365人学习下载读者可据此快速复现VIT分类流程并参考文档与演示材料整理实验报告与答辩内容。1. 从一次翻车说起为什么我用 ViT 重做 CIFAR-10 分类去年带本科生做深度学习大作业一个小组用 ResNet-18 在 CIFAR-10 上刷到 94% 准确率答辩时被问「为什么不用 Transformer」学生答「ViT 参数量太大小数据集训不动」。这个回答对了一半——原版 ViT 在 CIFAR-10 上从零训练确实会翻车准确率可能卡在 70% 出头但加上合适的数据增强、正则化和学习率调度纯 ViT 在 CIFAR-10 上做到 95% 以上是完全可以复现的。这篇笔记就围绕「Python ViT 实现 CIFAR-10 分类」这条主线把数据准备、模型搭建、训练调参、评估排错整条链路拆开讲清楚。适合正在做深度学习大作业的学生、想从 CNN 迁移到 Transformer 的工程师以及需要一份可复现基线代码的从业者。读完你能拿到一套能跑通的训练脚本、一组经过验证的超参配置以及一份踩坑清单。2. ViT 做 CIFAR-10 的选型逻辑与最小可跑通方案2.1 为什么原版 ViT 在小图上会「水土不服」ViT 的核心思路是把图像切成固定大小的 patch每个 patch 展平后过一个线性层变成 token再送进标准 Transformer Encoder。原版论文里 ViT-Base 用的 patch 大小是 16×16输入分辨率 224×224这样一张图有 196 个 token。CIFAR-10 的图像只有 32×32如果还用 16×16 的 patch一张图只能切出 4 个 token序列太短self-attention 根本学不到有意义的空间关系这是第一个坑。第二个问题是数据量。ViT 没有 CNN 那种平移不变性和局部归纳偏置它需要大量数据才能学到「相邻像素相关」这种先验。CIFAR-10 训练集只有 50000 张原版 ViT 从零训练会严重过拟合。常见做法有两种一是把 patch 调小到 4×4让 token 数变成 64序列长度够用二是引入强数据增强RandAugment、CutMix、MixUp和正则化DropPath、Label Smoothing把过拟合压下去。我一般会两个一起上。第三个问题是位置编码。patch 变小后 token 数变多可学习的位置编码参数量也跟着涨。CIFAR-10 这种小图用 4×4 patch 加可学习位置编码就够了不需要插值。2.2 环境准备与依赖安装先确认 Python 版本建议 3.8 以上。PyTorch 装 GPU 版本CUDA 版本按自己显卡驱动选。下面这套命令是我在 Ubuntu 20.04 RTX 3060 上验证过的# 创建虚拟环境避免污染系统 Python python -m venv vit_cifar source vit_cifar/bin/activate # 安装 PyTorchCUDA 11.8 版本按自己环境改 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 安装训练辅助库 pip install timm tensorboard tqdm numpy pillowtimm里有现成的 ViT 实现和预训练权重但做作业建议自己手写一遍模型结构理解更透。下面代码不依赖 timm 的模型定义只用了它的数据增强部分。2.3 数据加载与增强流水线CIFAR-10 用 torchvision 直接下载。训练集做 RandAugment CutMix测试集只做归一化。归一化的均值和方差用 CIFAR-10 的统计值import torch from torchvision import datasets, transforms from timm.data import RandAugment, Mixup # CIFAR-10 通道均值和标准差别用 ImageNet 的会掉点 CIFAR_MEAN (0.4914, 0.4822, 0.4465) CIFAR_STD (0.2470, 0.2435, 0.2616) train_transform transforms.Compose([ transforms.RandomCrop(32, padding4), # 先 padding 再随机裁剪 transforms.RandomHorizontalFlip(), # 水平翻转 RandAugment(num_ops2, magnitude9), # 自动增强2 个操作强度 9 transforms.ToTensor(), transforms.Normalize(CIFAR_MEAN, CIFAR_STD), transforms.RandomErasing(p0.25), # 随机擦除防过拟合 ]) test_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize(CIFAR_MEAN, CIFAR_STD), ]) train_set datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtrain_transform) test_set datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtest_transform) train_loader torch.utils.data.DataLoader(train_set, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue) test_loader torch.utils.data.DataLoader(test_set, batch_size256, shuffleFalse, num_workers4, pin_memoryTrue)参数说明RandomCrop(32, padding4)是 CIFAR-10 的标准操作先四周补 4 像素再裁回 32×32模拟平移。RandAugment的magnitude9是我试出来比较稳的值再高会破坏语义。RandomErasing(p0.25)概率别超过 0.3否则欠拟合。batch_size128在 8GB 显存上跑 ViT-Small 刚好显存不够就降到 64 并开梯度累积。2.4 手写 ViT 模型结构下面是一个适合 CIFAR-10 的 ViT 变体patch 大小 4嵌入维度 384深度 7注意力头数 6。这个配置参数量约 5.5M比 ResNet-18 还小import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim384): super().__init__() self.num_patches (img_size // patch_size) ** 2 # 64 个 patch self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # [B, embed_dim, H/P, W/P] x x.flatten(2).transpose(1, 2) # [B, num_patches, embed_dim] return x class ViTForCIFAR(nn.Module): def __init__(self, num_classes10, embed_dim384, depth7, num_heads6, mlp_ratio4.0, drop_path0.1): super().__init__() self.patch_embed PatchEmbed(embed_dimembed_dim) num_patches self.patch_embed.num_patches # 可学习位置编码比正弦编码更适合小数据集 self.pos_embed nn.Parameter(torch.zeros(1, num_patches, embed_dim)) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_drop nn.Dropout(0.1) # Transformer Encoder用 nn.TransformerEncoderLayer 堆叠 encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardint(embed_dim * mlp_ratio), dropout0.1, 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) # 初始化 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) def forward(self, x): B x.shape[0] x self.patch_embed(x) # [B, 64, 384] cls_tokens self.cls_token.expand(B, -1, -1) # [B, 1, 384] x torch.cat([cls_tokens, x], dim1) # [B, 65, 384] x x self.pos_embed[:, :x.size(1)] # 加位置编码 x self.pos_drop(x) x self.encoder(x) x self.norm(x[:, 0]) # 取 cls token return self.head(x)逻辑说明PatchEmbed用卷积实现切 patch比手动 reshape 快。cls_token是可学习的分类 token最后取它的输出做分类。位置编码用可学习参数初始化标准差 0.02。nn.TransformerEncoderLayer的batch_firstTrue让输入格式是[B, seq, dim]和 patch embed 输出对齐。drop_path参数这里没实际用上如果要加 DropPath 需要自己写或者用 timm 的DropPath模块替换残差连接。2.5 训练循环与关键超参训练用 AdamW学习率 3e-4权重衰减 0.05余弦退火加 5 轮 warmup。标签平滑 0.1。这些值是我在 CIFAR-10 上试了十几组后比较稳的import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR device torch.device(cuda if torch.cuda.is_available() else cpu) model ViTForCIFAR().to(device) # 标签平滑缓解过拟合 criterion nn.CrossEntropyLoss(label_smoothing0.1) # AdamWViT 对权重衰减敏感别用 SGD optimizer optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) # 5 轮 warmup 余弦退火总 100 轮 warmup LinearLR(optimizer, start_factor0.01, total_iters5) cosine CosineAnnealingLR(optimizer, T_max95, eta_min1e-6) scheduler SequentialLR(optimizer, schedulers[warmup, cosine], milestones[5]) for epoch in range(100): model.train() total_loss 0 for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() # 梯度裁剪防梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() scheduler.step() # 每 5 轮测一次 if (epoch 1) % 5 0: model.eval() correct total 0 with torch.no_grad(): for imgs, labels in test_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) _, pred outputs.max(1) correct (pred labels).sum().item() total labels.size(0) print(fEpoch {epoch1}, Loss {total_loss/len(train_loader):.4f}, Acc {correct/total:.4f})参数说明lr3e-4是 ViT 微调的常用起点从零训练可以再高一点到 1e-3但容易震荡。weight_decay0.05比 CNN 常用的 1e-4 大很多因为 Transformer 对过拟合更敏感。clip_grad_norm_的max_norm1.0是保险丝ViT 训练前期梯度偶尔会炸。label_smoothing0.1能提 0.5 到 1 个点。3. 训练过程中的参数调优与显存控制3.1 学习率和 batch size 的联动关系ViT 对学习率很敏感。我试过 lr1e-3 从零训练前 10 轮 loss 直接 NaN原因是注意力层的梯度在初始化时方差偏大。后来改成 3e-4 加 warmup 就稳了。batch size 和学习率要联动batch 翻倍lr 也大致翻倍。128 的 batch 配 3e-4256 的 batch 配 6e-4。如果显存不够只能跑 64 的 batchlr 降到 1.5e-4同时把 warmup 轮数加到 10。另一个经验是ViT 在 CIFAR-10 上前 20 轮准确率涨得很慢可能只有 60% 出头别急着调参继续跑。30 轮之后开始快速上升60 轮左右到 90%100 轮能到 95%。如果 50 轮还卡在 70%那大概率是数据增强太狠或者学习率不对。3.2 显存不够时的三种降级方案8GB 显存跑 ViT-Small batch 128 大概占 6.5GB如果同时开 TensorBoard 和多个 DataLoader worker可能爆。三种降级方案按优先级排第一降 batch size 到 64同时开梯度累积两步等效 batch 还是 128。代码改法是在 loss.backward() 前加loss loss / 2每两步 optimizer.step() 一次。第二用混合精度训练。PyTorch 的torch.cuda.amp能省 30% 到 40% 显存速度还快scaler torch.cuda.amp.GradScaler() for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(imgs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()注意clip_grad_norm_要放在scaler.unscale_之后否则裁剪的是缩放后的梯度数值不对。第三把模型深度从 7 降到 5嵌入维度从 384 降到 256。参数量降到 2.8M准确率大概掉 1 到 1.5 个点但显存占用减半。3.3 用 TensorBoard 盯住三个关键曲线训练时至少要看三条曲线训练 loss、测试准确率、学习率。训练 loss 震荡不降检查学习率和数据增强强度。测试准确率涨到某个点后回落说明过拟合加大 weight_decay 或 drop_path。学习率曲线如果 warmup 阶段就冲太高把start_factor从 0.01 降到 0.001。tensorboard --logdir./runs --port6006在代码里加SummaryWriter记录每轮写一次 loss 和 acc。这个习惯能帮你省下大量盲调时间。4. 避坑与排查CIFAR-10 上跑 ViT 的五个血泪教训4.1 准确率卡在 70% 上不去现象训练 50 轮测试准确率一直在 70% 附近晃loss 降得很慢。原因patch 大小设成了 1632×32 的图只切出 4 个 token序列太短注意力学不到东西。解决把 patch_size 改成 4token 数变成 64。如果显存够改成 2 更好token 数 256但计算量涨 4 倍。CIFAR-10 上 patch4 是性价比最高的。4.2 训练 loss 正常但测试准确率远低于训练准确率现象训练准确率 99%测试只有 85%差距 14 个点。原因过拟合。ViT 参数量相对 CIFAR-10 数据量还是偏大加上数据增强不够强。解决三件事一起做。weight_decay 从 0.05 提到 0.1RandAugment 的 magnitude 从 9 提到 12加 DropPath概率 0.1 到 0.2。如果还不行用 CutMix 替换 RandomErasingCutMix 的正则化效果更强。4.3 训练到一半 loss 突然变 NaN现象前 30 轮正常第 31 轮 loss 突然 NaN之后所有输出都是 NaN。原因梯度爆炸。ViT 的注意力层在某个 batch 上梯度范数突然冲高AdamW 更新后参数飞了。解决加梯度裁剪max_norm1.0。如果已经 NaN 了从上一个 checkpoint 恢复把学习率降一半。另外检查数据里有没有损坏的图片CIFAR-10 偶尔有下载不完整的样本用try/except跳过。4.4 显存溢出但 batch size 已经很小现象batch size 降到 32 还是 OOM。原因num_workers设太大每个 worker 都复制一份数据到显存或者测试时没加torch.no_grad()测试集前向也建了计算图。解决num_workers设 2 到 4 就够别超过 CPU 核数。测试循环必须包在with torch.no_grad():里。另外pin_memoryTrue在显存紧张时反而增加开销可以关掉。4.5 换用预训练权重后准确率反而下降现象加载了 ImageNet 预训练的 ViT 权重微调后测试准确率比从零训练还低。原因ImageNet 预训练用的 patch 是 16位置编码长度 196CIFAR-10 用 patch4 位置编码长度 64直接加载会形状不匹配。强行插值位置编码会破坏语义。解决要么把 patch 也设成 16但那样 CIFAR-10 上效果差要么只加载 patch embed 和 encoder 的权重位置编码重新初始化。我一般做作业直接从零训练CIFAR-10 数据量够 ViT-Small 收敛。5. 进阶技巧把 CIFAR-10 准确率推到 96% 以上5.1 用 CutMix 和 MixUp 的组合增强前面用的是 RandAugment RandomErasing想再往上提加 CutMix 和 MixUp。这两个都是样本混合策略CutMix 把一张图的部分区域替换成另一张图MixUp 把两张图按比例线性混合。代码实现import numpy as np def cutmix_data(x, y, alpha1.0): lam np.random.beta(alpha, alpha) batch_size x.size(0) index torch.randperm(batch_size).to(x.device) # 随机选一个矩形区域 bbx1, bby1, bbx2, bby2 rand_bbox(x.size(), lam) x[:, :, bbx1:bbx2, bby1:bby2] x[index, :, bbx1:bbx2, bby1:bby2] # 调整 lambda 为实际区域面积比 lam 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (x.size(-1) * x.size(-2))) return x, y, y[index], lam def rand_bbox(size, lam): W, H size[2], size[3] cut_rat np.sqrt(1. - lam) cut_w, cut_h int(W * cut_rat), int(H * cut_rat) cx, cy np.random.randint(W), np.random.randint(H) bbx1 np.clip(cx - cut_w // 2, 0, W) bby1 np.clip(cy - cut_h // 2, 0, H) bbx2 np.clip(cx cut_w // 2, 0, W) bby2 np.clip(cy cut_h // 2, 0, H) return bbx1, bby1, bbx2, bby2训练时 loss 要算两次按 lam 加权outputs model(imgs) loss lam * criterion(outputs, labels_a) (1 - lam) * criterion(outputs, labels_b)CutMix 和 MixUp 交替用每个 batch 随机选一种概率各 50%。这两个增强能把准确率从 95% 推到 96% 左右但训练轮数要加到 200 轮收敛更慢。5.2 用 EMA 权重做最终评估指数移动平均EMA是训练后期提点的利器。维护一份模型参数的滑动平均评估时用 EMA 权重而不是当前权重。实现很简单class EMA: def __init__(self, model, decay0.999): self.model model self.decay decay self.shadow {k: v.clone().detach() for k, v in model.state_dict().items()} def update(self): for k, v in self.model.state_dict().items(): if v.dtype.is_floating_point: self.shadow[k] self.decay * self.shadow[k] (1 - self.decay) * v else: self.shadow[k] v def apply(self): self.model.load_state_dict(self.shadow)每步训练后调ema.update()评估前调ema.apply()评估完再恢复原权重。decay 设 0.999训练轮数少的话设 0.99。EMA 能稳定提 0.3 到 0.5 个点而且几乎不增加训练开销。5.3 验证方法别只看最终准确率大作业答辩时老师常问「你怎么证明模型没作弊」。除了最终准确率至少还要看三个指标每类准确率、混淆矩阵、测试集上的 loss 曲线。CIFAR-10 里猫和狗容易混飞机和船也容易混混淆矩阵能暴露这些问题。如果某一类准确率特别低检查那一类的训练样本有没有被增强破坏。from sklearn.metrics import confusion_matrix import seaborn as sns model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in test_loader: imgs imgs.to(device) outputs model(imgs) _, pred outputs.max(1) all_preds.extend(pred.cpu().numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) sns.heatmap(cm, annotTrue, fmtd, cmapBlues)最后说个习惯我每次跑完实验会把配置文件、随机种子、最终准确率写进一个experiment_log.md下次调参直接翻记录比凭记忆靠谱。ViT 在 CIFAR-10 上不是玄学参数对了就能复现。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑