ARTICLE DETAIL

资讯详情

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

Python基于ViT实现CIFAR-10分类:从零训练与调参实战

Python基于ViT实现CIFAR-10分类:从零训练与调参实战 简介这份资源是面向深度学习课程大作业与计算机视觉入门者的完整实践包围绕Python与Vision TransformerVIT实现CAFIR10图像分类任务展开。内容从VIT将图像切分为patch、经自注意力机制捕获全局信息的原理出发结合CAFIR10数据集的加载、预处理与训练流程帮助读者理解Transformer架构在视觉分类中的落地方式适合需要完成课程设计或想从CNN过渡到VIT的学习者。压缩包共21个文件约11.25MB包含7个ipynb交互式笔记本用于分步实验、3个py脚本承载核心训练与推理逻辑、3份docx与3份pptx整理原理讲解和汇报材料另有txt说明、csv数据记录等辅助文件结构清晰便于按模块查阅。目前已有365人学习下载。读者可据此获得可运行的分类项目源码、配套文档与演示文稿并参考手写数字识别、机器翻译、LSTM自动写诗等同类作业案例快速搭建实验环境、复现训练过程并完成报告撰写。1. 用 ViT 啃下 CIFAR-10一份能跑通、能交作业的深度学习实战路线CIFAR-10 是深度学习入门绕不开的数据集10 个类别、6 万张 32×32 小图看着简单但真拿 ViTVision Transformer去做很多人第一次跑出来的准确率还不如一个三层 CNN。问题不在 ViT 不行而在于 ViT 天生是给 224×224 大图设计的直接怼到 32×32 上patch 切完只剩 64 个 token位置信息稀薄、归纳偏置几乎为零训练不收敛是常态。这份「Python 基于 ViT 实现 CIFAR-10 分类」的大作业核心要解决的就是怎么把 Transformer 那套自注意力机制在小分辨率图像上真正调通并且把代码、文档、实验记录整理成一份能交、能复现的完整工程。适合正在做课程大作业的学生也适合想从 CNN 转到 ViT 的工程师——下面这套流程我按自己带过几届学生的经验拆开讲参数、坑点、验证方法都给到。2. ViT 做 CIFAR-10 的原理与选型为什么不能直接套 ImageNet 那套2.1 ViT 的核心机制与 CIFAR-10 的尺寸矛盾ViT 的流程可以拆成四步切 patch、线性投影、加位置编码、送进 Transformer Encoder。以标准 ViT-Base 为例输入 224×224patch 大小 16×16得到 14×14196 个 patch每个 patch 展平成 768 维向量再加上一个 [CLS] token 和位置编码总共 197 个 token 进入 12 层 Encoder。自注意力在 197 个 token 之间做全局交互这是它比 CNN 强的地方——感受野从一开始就是全图。但 CIFAR-10 是 32×32。如果还用 patch16只能切出 2×24 个 patch序列长度 5含 CLSTransformer 根本没有足够的 token 去建模空间关系注意力图几乎退化成常数。这就是为什么直接套 ImageNet 预训练的 ViT 到 CIFAR-10 上微调后准确率经常卡在 70% 出头。解决办法有两个方向一是缩小 patch 尺寸比如 patch4得到 8×864 个 patch序列长度 65信息量够了二是用混合架构前面接一个轻量 CNN 做下采样或特征提取再把特征图切成 token。大作业里我一般推荐第一种纯 ViT 结构清晰代码量可控实验对比也好写。2.2 从零训练还是预训练微调两条路线的取舍做课程作业时间通常只有一两周算力也就是一张消费级显卡。这时候选型就很关键路线数据需求训练时间单卡预期准确率适合场景从零训练 ViT-small5 万训练样本2-4 小时85%-90%想讲清楚 ViT 原理实验完整ImageNet 预训练微调同上30 分钟-1 小时95%追求指标快速出结果CNNViT 混合同上1-2 小时90%-93%想对比不同架构从零训练的好处是你能完整展示数据增强、学习率调度、正则化这些技巧文档写起来有内容预训练微调则容易陷入「调参玄学」因为大部分能力来自预训练权重你自己的贡献不好体现。我一般建议主实验从零训练附一组预训练微调做对比这样既有深度又有说服力。2.3 数据增强与正则化小数据集上 ViT 的救命稻草ViT 没有 CNN 的平移不变性和局部性归纳偏置所以在小数据集上极易过拟合。CIFAR-10 只有 5 万张训练图必须靠增强把数据「撑大」。常用的组合是 RandomCrop(32, padding4) RandomHorizontalFlip CutMix 或 MixUp。CutMix 在 ViT 上效果尤其明显因为它强迫模型关注局部与全局的关系和自注意力的特性互补。正则化方面Dropout 放在 patch embedding 后和每个 Encoder block 里DropPath随机深度对深层 ViT 很关键weight decay 用 0.05 起步。这些参数不是拍脑袋后面章节会给具体配置。3. 环境搭建与数据管线把 CIFAR-10 喂给 ViT 的完整代码3.1 依赖安装与目录结构先确认 Python 版本建议 3.8 以上PyTorch 1.12。安装命令如下pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install timm tensorboard matplotlib numpytimm里有现成的 ViT 实现但大作业建议自己写一遍 Encoder否则答辩时说不清细节。目录结构按这个来vit-cifar10/ ├── data/ # 自动下载的 CIFAR-10 ├── models/vit.py # ViT 模型定义 ├── utils/dataset.py # 数据加载与增强 ├── train.py # 训练入口 ├── eval.py # 评估脚本 └── configs/base.yaml # 超参数配置3.2 数据增强管线的代码实现import torch from torchvision import datasets, transforms def build_transform(is_train, img_size32): if is_train: return transforms.Compose([ transforms.RandomCrop(img_size, padding4), transforms.RandomHorizontalFlip(), transforms.AutoAugment(transforms.AutoAugmentPolicy.CIFAR10), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) return transforms.Compose([ transforms.Resize(img_size), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_set datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformbuild_transform(True)) val_set datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformbuild_transform(False))这里有几个点要注意。归一化的均值方差是 CIFAR-10 的统计值不要用 ImageNet 的否则输入分布偏移会影响收敛。AutoAugment 的 CIFAR10 策略比手动组合增强更稳但会拖慢数据加载如果 GPU 利用率低于 70%可以把 num_workers 调到 4 或 8。CutMix 和 MixUp 不在 transform 里做而是在训练循环中按 batch 应用后面训练章节会写。3.3 ViT 模型定义patch embedding 与位置编码import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim192): super().__init__() self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # B, C, H/P, W/P x x.flatten(2).transpose(1, 2) # B, N, C return x class ViT(nn.Module): def __init__(self, img_size32, patch_size4, embed_dim192, depth6, num_heads6, mlp_ratio4, num_classes10, drop_rate0.1, drop_path_rate0.1): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, 3, embed_dim) 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, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(drop_rate) # Transformer Encoder 用 nn.TransformerEncoderLayer 堆叠 encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardint(embed_dim * mlp_ratio), dropoutdrop_rate, activationgelu, batch_firstTrue, norm_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) cls self.cls_token.expand(B, -1, -1) x torch.cat([cls, x], dim1) x x self.pos_embed x self.pos_drop(x) x self.encoder(x) x self.norm(x[:, 0]) # 取 CLS token return self.head(x)patch_size4 是关键选择32/48得到 64 个 patch序列长度 65计算量适中。embed_dim192、depth6、num_heads6 是 ViT-Small 的缩小版参数量约 5M单卡 8G 显存足够。norm_firstTrue 表示 Pre-LN训练更稳定这是小数据集上必须开的。位置编码用可学习参数初始化用 trunc_normal_不要用全零。4. 训练策略与调参让 ViT 在 CIFAR-10 上真正收敛4.1 优化器、学习率与 warmup 配置ViT 对优化器很敏感。AdamW 比 SGD 收敛快但最终精度可能略低如果训练轮数够200 epoch 以上SGD momentum 0.9 往往更好。大作业时间有限我一般用 AdamWbetas(0.9, 0.999)weight decay 0.05学习率 3e-4 起步配合 cosine 退火和 5 个 epoch 的 warmup。import torch.optim as optim from torch.optim.lr_scheduler import LambdaLR import math def build_optimizer(model, lr3e-4, wd0.05): decay, no_decay [], [] for name, param in model.named_parameters(): if not param.requires_grad: continue if param.ndim 1 or name.endswith(.bias): no_decay.append(param) else: decay.append(param) return optim.AdamW([ {params: decay, weight_decay: wd}, {params: no_decay, weight_decay: 0.0} ], lrlr, betas(0.9, 0.999)) def warmup_cosine(optimizer, warmup_epochs, total_epochs, steps_per_epoch): def fn(step): if step warmup_epochs * steps_per_epoch: return step / (warmup_epochs * steps_per_epoch) progress (step - warmup_epochs * steps_per_epoch) / \ ((total_epochs - warmup_epochs) * steps_per_epoch) return 0.5 * (1 math.cos(math.pi * progress)) return LambdaLR(optimizer, fn)weight decay 不作用于 LayerNorm 和 bias这是 ViT 训练的常识否则会抑制归一化层的缩放能力。warmup 不能省ViT 初期梯度方差大直接上大学习率容易炸。4.2 CutMix 与 MixUp 在训练循环中的实现import numpy as np def cutmix_data(x, y, alpha1.0): lam np.random.beta(alpha, alpha) index torch.randperm(x.size(0)).to(x.device) bbx1, bby1, bbx2, bby2 rand_bbox(x.size(), lam) x[:, :, bbx1:bbx2, bby1:bby2] x[index, :, bbx1:bbx2, bby1:bby2] 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训练时按概率切换 CutMix 和 MixUp比如各 50%或者只用 CutMix。损失函数用软标签交叉熵criterion nn.CrossEntropyLoss() # 前向 if use_cutmix: x, y_a, y_b, lam cutmix_data(x, y) out model(x) loss lam * criterion(out, y_a) (1 - lam) * criterion(out, y_b) else: out model(x) loss criterion(out, y)CutMix 的 alpha 设 1.0MixUp 也设 1.0这是 CIFAR 上的常用值。注意 CutMix 后要重新计算 lam因为裁剪区域可能被边界截断。4.3 训练轮数、batch size 与显存权衡CIFAR-10 上从零训练 ViT-small建议 200-300 epochbatch size 128 或 256。batch 太小如 32会导致 BatchNorm 统计不稳但 ViT 用 LayerNorm所以 batch 64 也能跑只是训练更慢。显存不够时用梯度累积模拟大 batchaccum_steps 4 for i, (x, y) in enumerate(loader): out model(x) loss criterion(out, y) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()梯度累积时学习率要按实际 batch 缩放否则等效学习率变了。另外混合精度训练AMP能省 30%-40% 显存但 ViT 的 LayerNorm 在 fp16 下容易溢出建议用 bfloat16 或保持 fp32 的 LayerNorm。5. 避坑与排查ViT 训 CIFAR-10 最常见的 5 个翻车现场5.1 损失不下降准确率卡在 10%现象训练几个 epoch 后 loss 几乎不变准确率等于随机猜。 原因patch_size 太大导致序列太短或者位置编码初始化有问题。 解决确认 patch_size ≤ 8pos_embed 用 trunc_normal_(std0.02) 初始化不要全零。另外检查学习率是否过大warmup 是否生效。5.2 训练准确率很高验证准确率低 20 个点现象train acc 99%val acc 75%典型过拟合。 原因数据增强太弱或者 drop_path_rate 设得太小。 解决加上 CutMix/MixUpdrop_path_rate 从 0.1 提到 0.2weight decay 从 0.05 提到 0.1。如果还不行减少模型深度或 embed_dim。5.3 训练中途 loss 突然变成 NaN现象前几十个 epoch 正常突然 lossnan。 原因混合精度下 LayerNorm 溢出或者学习率在 warmup 后跳变。 解决改用 bfloat16或者在 LayerNorm 前强制转 fp32。检查 scheduler 是否在 warmup 结束后直接跳到峰值学习率cosine 退火应该平滑过渡。5.4 GPU 利用率低训练速度慢现象nvidia-smi 显示 GPU 利用率只有 30%-50%。 原因数据加载是瓶颈num_workers 太少或增强太复杂。 解决num_workers 设为 CPU 核数的一半pin_memoryTruepersistent_workersTrue。AutoAugment 比较慢可以换成 RandAugment 或手动增强。5.5 复现结果和论文/基线对不上现象同样的配置别人跑 90%你跑 85%。 原因随机种子、数据顺序、增强策略的细微差异。 解决固定所有随机种子torch、numpy、random用 deterministic 模式记录完整的配置文件。ViT 对初始化敏感多跑几个种子取平均更可靠。6. 进阶技巧用注意力可视化和消融实验把大作业做出深度6.1 注意力图可视化验证 ViT 到底学到了什么大作业如果只报一个准确率很难拿高分。加一组注意力可视化能直接说明模型关注了哪些区域。方法很简单取最后一层 Encoder 的注意力权重对 CLS token 那一行做 reshape得到 8×8 的注意力图再上采样到 32×32 叠加在原图上。def get_attention_map(model, x): # 注册 hook 抓取最后一层注意力 attn_maps [] def hook(module, input, output): attn_maps.append(output) handle model.encoder.layers[-1].self_attn.register_forward_hook(hook) _ model(x) handle.remove() # attn: B, num_heads, N, N attn attn_maps[0] cls_attn attn[:, :, 0, 1:] # CLS 对 patch 的注意力 cls_attn cls_attn.mean(dim1) # 多头平均 return cls_attn.reshape(-1, 8, 8)可视化时选几张正确分类和错误分类的图对比正确分类的注意力通常集中在目标物体上错误分类的注意力往往散在背景。这个分析写进文档比单纯堆准确率有说服力得多。6.2 消融实验设计patch size、增强、正则化的影响消融实验不用多三组就够patch size 取 2/4/8数据增强取「无/基础/CutMix」drop_path_rate 取 0/0.1/0.2。每组跑 100 epoch记录 val acc。表格一列结论自然出来。我自己的经验是 patch4 比 patch8 高 3-5 个点CutMix 比基础增强高 2-3 个点drop_path 从 0 到 0.1 提升最明显再到 0.2 就趋于平缓。6.3 模型导出与推理脚本最后交作业前把最好的权重导出写一个干净的推理脚本支持单张图片预测。用 torch.jit.trace 或 ONNX 导出都行ONNX 更通用dummy torch.randn(1, 3, 32, 32) torch.onnx.export(model, dummy, vit_cifar10.onnx, input_names[input], output_names[logits], dynamic_axes{input: {0: batch}})推理时记得用验证集的 transform不要用训练增强。CIFAR-10 的类别名按顺序写死airplane, automobile, bird, cat, deer, dog, frog, horse, ship, truck。这套流程我带学生跑过好几轮最深的教训是ViT 在小图上的调参空间比 CNN 窄patch size 和 warmup 这两个参数一旦设错后面怎么调都救不回来。所以别急着堆 epoch先把 patch4、warmup5、AdamW 这三点固定住再动其他参数。希望帮到你。本文还有配套的精品资源点击获取
返回列表