news 2026/10/5 4:09:11

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Python+ViT实战:CIFAR-10图像分类从70%到96%的调优指南

简介:这份资源面向深度学习初学者与课程实践者,提供一套基于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 # 安装 PyTorch,CUDA 11.8 版本,按自己环境改 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 安装训练辅助库 pip install timm tensorboard tqdm numpy pillow

timm里有现成的 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, padding=4), # 先 padding 再随机裁剪 transforms.RandomHorizontalFlip(), # 水平翻转 RandAugment(num_ops=2, magnitude=9), # 自动增强,2 个操作,强度 9 transforms.ToTensor(), transforms.Normalize(CIFAR_MEAN, CIFAR_STD), transforms.RandomErasing(p=0.25), # 随机擦除,防过拟合 ]) test_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(CIFAR_MEAN, CIFAR_STD), ]) train_set = datasets.CIFAR10(root='./data', train=True, download=True, transform=train_transform) test_set = datasets.CIFAR10(root='./data', train=False, download=True, transform=test_transform) train_loader = torch.utils.data.DataLoader(train_set, batch_size=128, shuffle=True, num_workers=4, pin_memory=True) test_loader = torch.utils.data.DataLoader(test_set, batch_size=256, shuffle=False, num_workers=4, pin_memory=True)

参数说明:RandomCrop(32, padding=4)是 CIFAR-10 的标准操作,先四周补 4 像素再裁回 32×32,模拟平移。RandAugment的magnitude=9是我试出来比较稳的值,再高会破坏语义。RandomErasing(p=0.25)概率别超过 0.3,否则欠拟合。batch_size=128在 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_size=32, patch_size=4, in_chans=3, embed_dim=384): super().__init__() self.num_patches = (img_size // patch_size) ** 2 # 64 个 patch self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_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_classes=10, embed_dim=384, depth=7, num_heads=6, mlp_ratio=4.0, drop_path=0.1): super().__init__() self.patch_embed = PatchEmbed(embed_dim=embed_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_model=embed_dim, nhead=num_heads, dim_feedforward=int(embed_dim * mlp_ratio), dropout=0.1, activation='gelu', batch_first=True ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) # 初始化 nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.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], dim=1) # [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_first=True让输入格式是[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_smoothing=0.1) # AdamW,ViT 对权重衰减敏感,别用 SGD optimizer = optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.05) # 5 轮 warmup + 余弦退火,总 100 轮 warmup = LinearLR(optimizer, start_factor=0.01, total_iters=5) cosine = CosineAnnealingLR(optimizer, T_max=95, eta_min=1e-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_norm=1.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(f'Epoch {epoch+1}, Loss {total_loss/len(train_loader):.4f}, Acc {correct/total:.4f}')

参数说明:lr=3e-4是 ViT 微调的常用起点,从零训练可以再高一点到 1e-3,但容易震荡。weight_decay=0.05比 CNN 常用的 1e-4 大很多,因为 Transformer 对过拟合更敏感。clip_grad_norm_的max_norm=1.0是保险丝,ViT 训练前期梯度偶尔会炸。label_smoothing=0.1能提 0.5 到 1 个点。

3. 训练过程中的参数调优与显存控制

3.1 学习率和 batch size 的联动关系

ViT 对学习率很敏感。我试过 lr=1e-3 从零训练,前 10 轮 loss 直接 NaN,原因是注意力层的梯度在初始化时方差偏大。后来改成 3e-4 加 warmup 就稳了。batch size 和学习率要联动:batch 翻倍,lr 也大致翻倍。128 的 batch 配 3e-4,256 的 batch 配 6e-4。如果显存不够只能跑 64 的 batch,lr 降到 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_norm=1.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 --port=6006

在代码里加SummaryWriter记录,每轮写一次 loss 和 acc。这个习惯能帮你省下大量盲调时间。

4. 避坑与排查:CIFAR-10 上跑 ViT 的五个血泪教训

4.1 准确率卡在 70% 上不去

现象:训练 50 轮,测试准确率一直在 70% 附近晃,loss 降得很慢。

原因:patch 大小设成了 16,32×32 的图只切出 4 个 token,序列太短,注意力学不到东西。

解决:把 patch_size 改成 4,token 数变成 64。如果显存够,改成 2 更好,token 数 256,但计算量涨 4 倍。CIFAR-10 上 patch=4 是性价比最高的。

4.2 训练 loss 正常但测试准确率远低于训练准确率

现象:训练准确率 99%,测试只有 85%,差距 14 个点。

原因:过拟合。ViT 参数量相对 CIFAR-10 数据量还是偏大,加上数据增强不够强。

解决:三件事一起做。weight_decay 从 0.05 提到 0.1;RandAugment 的 magnitude 从 9 提到 12;加 DropPath,概率 0.1 到 0.2。如果还不行,用 CutMix 替换 RandomErasing,CutMix 的正则化效果更强。

4.3 训练到一半 loss 突然变 NaN

现象:前 30 轮正常,第 31 轮 loss 突然 NaN,之后所有输出都是 NaN。

原因:梯度爆炸。ViT 的注意力层在某个 batch 上梯度范数突然冲高,AdamW 更新后参数飞了。

解决:加梯度裁剪,max_norm=1.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_memory=True在显存紧张时反而增加开销,可以关掉。

4.5 换用预训练权重后准确率反而下降

现象:加载了 ImageNet 预训练的 ViT 权重,微调后测试准确率比从零训练还低。

原因:ImageNet 预训练用的 patch 是 16,位置编码长度 196,CIFAR-10 用 patch=4 位置编码长度 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, alpha=1.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, decay=0.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, annot=True, fmt='d', cmap='Blues')

最后说个习惯:我每次跑完实验,会把配置文件、随机种子、最终准确率写进一个experiment_log.md,下次调参直接翻记录,比凭记忆靠谱。ViT 在 CIFAR-10 上不是玄学,参数对了就能复现。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/5 4:09:10

C++设计模式实战:用智能指针与RAII重构经典模式

在C里写设计模式,和你在博客或教材里看到的UML/Java示例完全是两码事。我见过太多人(包括我自己当年)把《Head First 设计模式》里的Java代码一行行翻译成C,结果遇到析构顺序崩掉、拷贝构造多复制一份资源、并发下单例被反复初始化…

作者头像 李华
网站建设 2026/10/5 4:09:08

OpenShell 可编程外壳框架:插件化架构、上下文隔离与命令审计实战

1. 从零认识 OpenShell:它到底解决什么问题第一次听到 OpenShell 这个名字,很多人会下意识以为它又是一个“终端美化工具”或者“命令行增强插件”。我最初也是这么想的,直到真正把它拉进项目里跑了一遍,才发现它的定位比想象中要…

作者头像 李华
网站建设 2026/10/5 4:08:21

Python深度学习手语翻译课设实战:从数据采集到模型推理全流程解析

简介:这份资源是大二期末课程设计项目,主题为基于Python与深度学习的手语翻译程序开发,面向计算机、人工智能、通信工程、自动化等专业的高校学生与教师,可用于课程设计、毕业设计、作业提交或项目初期立项演示,也适合…

作者头像 李华
网站建设 2026/10/5 4:08:10

一张 AI 冰山图,画出了三类人:用别人的、借别人的、造自己的

摘要流传很广的那张《How People See AI》冰山图,大多数人当成工具清单收藏了。但图里真正有价值的信息不在那 17 个工具名字上,而在它把人分成了三类:水面之上是消费者,水面第一层是流程设计者,水面最底部是基础设施搭…

作者头像 李华
网站建设 2026/10/5 4:08:02

弱电网下LCL-VSC次/超同步谐振的阻抗建模与Nyquist判据仿真复现

做并网变流器的人应该都有体会:电网一“弱”,各种低频振荡问题就全冒出来了。这里说的“弱电网”并不是电压不够高或者容量小,而是从变流器端口看进去,电网等效阻抗已经不能忽略,通常用短路比SCR来衡量。SCR越小&#…

作者头像 李华
网站建设 2026/10/5 4:07:57

弱电网下LCL-VSC阻抗建模与次/超同步谐振Nyquist分析

做风电、光伏并网或者储能变流器控制的同学,应该都听说过这么一句话:变流器电流环、电压环参数在理想电网下调得好好的,一接到弱电网就出幺蛾子——并网点电压畸变、电流里冒出低频振荡分量,严重的时候直接触发保护跳闸。前几年我…

作者头像 李华