尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch实现ResNet18图像分类:CIFAR10实战与调参全记录

PyTorch实现ResNet18图像分类:CIFAR10实战与调参全记录 简介基于PyTorch的ResNet18图像识别项目聚焦CIFAR10十类小图分类任务。资源面向深度学习和计算机视觉初学者也适合需要快速搭建残差网络基准实验的研究者完整呈现从数据预处理、模型构建到训练评估的工程流程。压缩包共26个文件约341MB包含Python训练与测试脚本、训练好的模型权重、CIFAR10原始数据分片、数据说明文档及训练过程可视化图片等可支撑代码阅读、结果复现与二次开发。已有1382人学习下载。亮点在于工程细节完整数据集加载与增强、残差块定义、损失函数与优化器配置均有对应实现还附带了模型预测效果图与曲线图便于对照检查训练效果借助权重文件可跳过训练直接进行推理也可在现有基础上调整超参数继续迭代是理解ResNet结构与PyTorch实战的实用样本。 不管是刚入门深度学习还是准备搞论文实验ResNet18加上CIFAR10这套组合都算得上是最经典的开胃菜。我最近整理资料时正好翻到一个命名为“ResNet18_CIFAR10.rar”的压缩包里面是之前跑通的一套训练代码和实验记录。借着这个契机我把整个项目的关键点、踩坑记录和可复用的操作流程完整梳理一遍希望对正在配环境或者卡在模型收敛的同学有点实际帮助。先说结论这套项目做的事情非常简单直接——用PyTorch搭建ResNet18网络在CIFAR10数据集上完成图像分类训练最终测试准确率能稳定达到92%~94%左右。它不能刷SOTA但特别适合用来理解残差网络的工作机制、掌握PyTorch的标准训练流程以及排查训练中的常见问题。如果你是第一次接触图像分类或者想找一个干净的基线模型做对比实验这份配置就是很好的起步模板。1. 项目核心思路拆解为什么偏偏是ResNet18和CIFAR101.1 CIFAR10数据集的特性与意义CIFAR10数据集包含60000张32x32像素的彩色图片分属10个类别每类6000张其中50000张用于训练10000张用于测试。这十个类别包括飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车都是日常生活里能接触到的物体。图像分辨率极低只有32x32人类都不一定能一眼辨认清楚这对模型而言反而是一种考验网络必须在极小的空间分辨率下提取足够有判别力的特征。这个数据集的优势在于它足够小单张图片才3KB左右整个数据集下载下来不到200MB训练一轮在普通消费级显卡上也就几十秒。我做这个项目时用的还是GTX 1660 Super训练100轮大概花了不到两个小时完全不心疼电费。同时CIFAR10又有一定的挑战性不像MNIST那样简单到随便一个线性模型就能到90%也不像ImageNet那样大到个人玩家根本跑不动。它处在“刚好能让你感受到调参乐趣”的甜点上。1.2 ResNet18的结构优势与残差思想ResNet18是残差网络系列里最轻量级的成员。它的核心创新在于引入了残差块通过跳跃连接把输入直接加到卷积层的输出上。换句话说每个残差块学习的不是完整的映射而是输入和输出之间的残差。这个设计的深层逻辑是当网络加深时如果恒等映射是最优解那么残差块只需要把权重逼近零即可这样梯度就能顺畅地回传不会因为层数增加而出现退化问题。放在实际项目中ResNet18一共有8个基础残差块加上开头一个7x7卷积和最后全连接层参数量大约1120万。这个体量对CIFAR10来说完全够用而且不会像ResNet50那样动辄几千万参数容易在小型数据集上过拟合。更关键的是ResNet18的训练速度适中显存占用不到2GB这意味着很多没有高端显卡的同学也能顺利跑通。1.3 技术选型的现实考虑我记得第一次做这个项目时脑子里转过的方案其实不止一个。有人推荐VGG16说结构简单容易理解但VGG16光是全连接层的参数量就占了大部分对CIFAR10这种小图来说严重浪费算力。也有人建议直接上ResNet50认为越深越强但32x32的输入根本喂不饱50层网络的感受野反而是ResNet18更匹配这个输入尺寸。所以最终选择ResNet18不是因为它最强而是因为它最合适。这种“匹配数据规模”的思路在水论文做实验时同样重要。一个基线模型如果本身就过拟合或者欠拟合后面你做的所有改进都无法得出可靠结论。2. 环境配置与数据准备从零跑通PyTorch环境2.1 软硬件环境清单这个项目的环境要求非常亲民。我用的是下面这套组合供参考操作系统Ubuntu 20.04Windows也能跑但Linux下调试更顺手GPUNVIDIA GTX 1660 Super 6GB显存CPUIntel i5-10400内存16GBPython版本3.8.10PyTorch版本1.10.1CUDA 11.3torchvision版本0.11.1如果电脑没有NVIDIA显卡用CPU跑也能出结果只是速度会慢很多100轮可能需要跑几个小时。这里建议优先使用conda创建独立环境避免把系统Python搞乱。创建命令很简单conda create -n resnet18 python3.8 conda activate resnet18 pip install torch1.10.1cu113 torchvision0.11.1cu113 -f https://download.pytorch.org/whl/torch_stable.html需要注意PyTorch的版本和CUDA版本要严格对应。如果你用的是更新版本的PyTorch去官方网站选择对应的安装命令即可。千万别直接pip install torch那样默认装CPU版本后面跑起来会发现训练速度慢到怀疑人生。2.2 CIFAR10数据集下载与本地化处理使用torchvision加载CIFAR10非常便捷核心代码就几行import torchvision.transforms as transforms from torchvision.datasets import CIFAR10 from torch.utils.data import DataLoader transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) train_set CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) test_set CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_set, batch_size128, shuffleTrue, num_workers4) test_loader DataLoader(test_set, batch_size128, shuffleFalse, num_workers4)这段代码里最容易忽略的是Normalize的均值方差。很多初学者会用ImageNet的均值方差(0.485, 0.456, 0.406)来归一化CIFAR10这其实是错误的。CIFAR10有自己的数据集统计特征上面代码里的数值是官方推荐的标准值。归一化不对训练前期loss下降会特别慢因为输入数据的分布没有被拉回标准正态分布梯度更新方向会被带偏。另外训练集和测试集的transform必须不同。训练集需要RandomCrop和RandomHorizontalFlip做数据增强测试集只用ToTensor和Normalize保证评估时数据的确定性。2.3 数据增强策略的选择数据增强是提升模型泛化能力极其有效的工具。在CIFAR10上我使用的RandomCrop(32, padding4)含义是先在原始图片四周填充4个像素的0然后在填充后的图片上随机裁剪回32x32。这样相当于给模型提供了一些平移扰动。RandomHorizontalFlip则是随机水平翻转利用了图片的对称性让模型不依赖于物体的朝向。我前后对比过加了这两个增强后测试准确率可以提高2~3个百分点。有人还会加CutOut或MixUp但这类高级增强会改变数据分布在CIFAR10这种数据量不算特别大的场景下副作用不好控制。作为基线项目保持简单稳定的增强策略就够了。3. 模型搭建与训练参数配置3.1 手动搭建ResNet18模型结构torchvision里虽然直接提供了现成的resnet18接口但那是针对ImageNet设计的输入是224x224第一个卷积层核大小是7x7stride2。如果原封不动用在CIFAR10的32x32图上会直接丢失大量边缘信息。因此推荐手动改造把第一个卷积层改成3x3、stride1、padding1并去掉开头的最大池化层。核心实现如下import torch import torch.nn as nn class BasicBlock(nn.Module): expansion 1 def __init__(self, in_channels, out_channels, stride1): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity self.shortcut(x) out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out identity out self.relu(out) return out class ResNet18(nn.Module): def __init__(self, num_classes10): super().__init__() self.in_channels 64 self.conv1 nn.Conv2d(3, 64, kernel_size3, stride1, padding1, biasFalse) self.bn1 nn.BatchNorm2d(64) self.relu nn.ReLU(inplaceTrue) self.layer1 self._make_layer(64, 2, stride1) self.layer2 self._make_layer(128, 2, stride2) self.layer3 self._make_layer(256, 2, stride2) self.layer4 self._make_layer(512, 2, stride2) self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.fc nn.Linear(512, num_classes) def _make_layer(self, out_channels, num_blocks, stride): strides [stride] [1] * (num_blocks - 1) layers [] for s in strides: layers.append(BasicBlock(self.in_channels, out_channels, strides)) self.in_channels out_channels return nn.Sequential(*layers) def forward(self, x): x self.conv1(x) x self.bn1(x) x self.relu(x) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) x self.avgpool(x) x torch.flatten(x, 1) x self.fc(x) return x这里有个细节很容易被忽略shortcut连接处的处理。当残差块的输入通道数和输出通道数不一致或者需要降低空间分辨率时stride2直接相加会报维度错误。因此必须通过1x1卷积调整通道数和尺寸。我在代码里用nn.Sequential()做了一个判断初始为空只有满足条件时才填入卷积和BN层这个写法非常清晰推荐照抄。3.2 损失函数与优化器的选择理由图像分类任务的标配是交叉熵损失函数nn.CrossEntropyLoss()。PyTorch的这个函数内部把LogSoftmax和NLLLoss合并了所以模型的最后一层不需要再手动加Softmax。很多人会画蛇添足地在网络输出后加Softmax再算loss结果训练时数值不稳定其实Softmax只用在推理阶段展示概率时损失计算直接用原始logits即可。优化器我选择的是SGD带动量而不是Adam。这不是我拍脑袋决定的而是基于大量实验观察。Adam收敛快但容易泛化性能差在CIFAR10这种中小型数据集上SGD配合CosineAnnealing学习率调度器往往能获得更高的最终精度。criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.1, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100)这里学习率设置0.1是因为SGDMomentum对学习率敏感且配合warmup策略从0.1起步可以在前几个epoch把模型从随机状态拉入可收敛区间。weight_decay设为5e-4等价于L2正则化能有效抑制过拟合。3.3 训练循环的标准写法训练循环本身并不复杂核心框架如下def train_one_epoch(model, train_loader, optimizer, criterion, device): model.train() total_loss 0 correct 0 total 0 for inputs, targets in train_loader: inputs, targets inputs.to(device), targets.to(device) optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, targets) loss.backward() optimizer.step() total_loss loss.item() * inputs.size(0) _, predicted outputs.max(1) total targets.size(0) correct predicted.eq(targets).sum().item() return total_loss / total, 100.0 * correct / total def evaluate(model, test_loader, criterion, device): model.eval() total_loss 0 correct 0 total 0 with torch.no_grad(): for inputs, targets in test_loader: inputs, targets inputs.to(device), targets.to(device) outputs model(inputs) loss criterion(outputs, targets) total_loss loss.item() * inputs.size(0) _, predicted outputs.max(1) total targets.size(0) correct predicted.eq(targets).sum().item() return total_loss / total, 100.0 * correct / total两个函数一定要区分model.train()和model.eval()。因为BatchNorm层在训练和推理时行为不同训练时使用当前batch的均值和方差计算归一化推理时使用训练阶段累计的全局统计量。如果漏掉eval()测试时BN层还在用当前batch的统计量结果会非常不稳定甚至出现准确率忽高忽低的情况这是新手最容易踩的坑。4. 训练细节与调参实战4.1 超参数组合与运行效果我最终选择了一套经过反复验证的训练配置完整列在下面超参数取值说明Batch Size128显存允许范围内尽可能大梯度更稳定Epochs100配合余弦退火足够收敛Initial LR0.1需要配合warmup直接全量训练会震荡Momentum0.9标准设置Weight Decay5e-4抑制过拟合Warmup Epochs5前5轮从0.01线性升到0.1按照这套配置训练前10轮训练准确率就爬到了约90%但测试准确率只有80%左右这是正常现象。到第50轮左右测试准确率能到90%最后100轮结束时测试准确率稳定在92.5%到93.8%之间。我特意尝试过不设置Weight Decay测试准确率掉到了91%左右把学习率从0.1换成0.01训练速度变慢最终精度也低了1.5个百分点。这说明调参不是玄学每个参数都有它存在的意义但也要结合数据规模和模型结构来综合判断。4.2 训练过程的Loss曲线观察从loss曲线能读到非常多的信息。一开始训练loss在1.8到2.0左右因为10分类任务随机初始状态的交叉熵大约是ln(10)2.30。前几个epoch loss快速下降这是模型从随机特征快速收敛到可辨识特征的过程。大约20轮后训练loss下降到0.5附近测试集loss也同步下降。但是80轮之后如果发现训练loss继续下降而测试loss开始回升那就是过拟合信号此时应该提前停止或者增强正则化。我把训练日志输出到控制台每轮打印一次当前epoch、学习率、训练loss、训练准确率、测试loss和测试准确率。看起来繁琐但对排查问题帮助极大特别是观察学习率是否按照余弦曲线衰减到接近零。4.3 常见问题与排查技巧实录很多人跑这个项目时会遇到“loss变成NaN”“测试准确率一直50%左右”等诡异现象。我把自己遇到过的以及帮助别人解决过的问题整理成了速查表问题现象可能原因解决方案Loss为NaN学习率过大梯度爆炸降低学习率到0.01以下或使用梯度裁剪训练准确率上不去数据归一化错误或网络结构错误检查Normalize均值和方差检查最后全连接层类别数测试准确率远低于训练过拟合增加数据增强、增大Weight Decay或使用DropoutBN层不收敛Batch Size太小导致统计量不准增大Batch Size到64以上或换用GroupNorm多GPU训练时报错设备ID不匹配/数据加载不均衡检查DataLoader的num_workers显存不足时减小Batch Size测试时每次结果波动大忘记调用model.eval()确认测试模式下BN层使用全局统计量这里特别想展开说两点。第一出现NaN时不要慌先检查学习率这是90%的情况。SGD在lr0.1时就可能出现梯度爆炸特别是数据集包含异常样本时。第二CIFAR10的类别顺序是固定的如果你自定义了数据集但标签顺序没对齐准确率就会一直卡在10%或50%这种问题排查起来特别隐蔽建议用tensorboard或脚本可视化一批样本检查标签是否匹配。4.4 模型保存与加载的坑训练结束后需要保存模型但很多人会直接把整个模型torch.save(model)这虽然方便但会留下隐患。推荐只保存state_dicttorch.save(model.state_dict(), resnet18_cifar10.pth)加载时要先构建模型结构再加载参数model ResNet18(num_classes10) model.load_state_dict(torch.load(resnet18_cifar10.pth)) model.eval()我踩过一个特别坑的细节跨机器加载模型时如果机器名、相对路径或者环境变量改变了经常会出现“size mismatch for fc.weight”之类的错误。这通常是因为模型结构不一致比如之前训练时类别数多了几个。解决办法是打印模型的state_dict键值对尺寸进行比对而不是盲目加strictFalse强制加载。5. 如何基于这个基线快速扩展跑通这套ResNet18_CIFAR10项目只是第一步。它更大的价值在于作为一个可靠基线可以支持后续一系列实验。我个人常做的扩展包括替换不同深度的残差网络把ResNet18换成ResNet34、ResNet50只需改动_make_layer的层数配置就能对比深度对准确率的影响。修改优化器在相同数据增强下对比SGD、AdamW、RMSProp的收敛速度与最终精度很多论文消融实验就是这么做的。做可视化分析用Grad-CAM输出模型关注的图像区域直观查看残差网络是否学到了符合直觉的特征。加入模型量化用PyTorch的量化工具把训练好的模型转为int8测试在CPU上的推理速度与精度损失这对边缘设备部署很有意义。当你熟悉了代码中的每个模块再进行这些扩展时会发现一切都很顺畅因为这套系统至少保证了数据加载、训练、评估、保存加载四个环节真实可靠。6. 经验总结与实用建议在我做过的所有图像分类项目中ResNet18_CIFAR10这套组合虽然简单但它像一块完美的练兵场。它帮你把“理解模型结构”“掌握训练流程”“排查训练问题”这三件事一次性打通。尤其是残差块里那个跳跃连接理解它的原理后你再去看各种现代结构比如DenseNet、Transformer的残差设计都会有一种豁然开朗的感觉。训练过程中我个人体会最深的一点是不要急着修改网络结构或堆复杂trick先把Baseline稳定跑出来确保代码流程没有任何隐藏bug。很多同学一上来就加各种注意力机制结果模型不收敛最后发现只是数据集归一化错了白折腾一周。先复现一个已知精度的模型再谈改进这是最稳妥也最高效的路线。最后再分享一个小技巧每次训练前手动设置随机种子包括Python、NumPy和PyTorch三个层面的种子这样能保证每次实验在相同环境下可复现。我之前用固定种子的代码在同等条件下跑了三次准确率偏差不超过0.2%这对写论文做对比实验非常重要毕竟任何一次实验的偶然波动都可能误导你的研究方向。希望这篇拆解能帮你又快又稳地把ResNet18跑在CIFAR10上并真正理解它背后的原理。本文还有配套的精品资源点击获取
返回列表