ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PyTorch实战CIFAR-100图像分类:CNN、ResNet与DenseNet多模型对比

PyTorch实战CIFAR-100图像分类:CNN、ResNet与DenseNet多模型对比 简介这是一份面向深度学习初学者与计算机视觉实践者的PyTorch图像分类实战资源聚焦CIFAR-100细粒度分类任务提供多种主流网络架构的完整可运行实现。资源涵盖ResNet、DenseNet、MobileNetV2、ShuffleNetV2、SENet、WideResNet、Inception系列、NASNet等20余种模型辅以训练train.py、测试test.py、数据加载dataset.py、工具函数utils.py及学习率查找lr_finder.py等模块结构清晰、即插即用。压缩包共28个文件含26个Python源码文件承担模型定义、训练逻辑与评估功能、1份README说明文档和1个.gitignore配置文件总大小仅43KB轻量高效。已有1725人下载学习适合课程实验、模型对比研究或竞赛基线搭建代码注释充分、模块解耦良好便于理解各算法在CIFAR-100上的适配细节与性能差异。1. 项目概述与核心价值做图像分类这件事说难不难说简单也真不简单。CIFAR-100作为计算机视觉领域最经典的基准数据集之一几乎是每个深度学习从业者绕不开的练手项目。我这次基于PyTorch实现了一个完整的CIFAR-100分类项目并且在同一套代码框架里集成了多种主流算法从最基础的自定义CNN到ResNet、DenseNet这类现代骨干网络全部跑通并做了对比实验。为什么选CIFAR-100而不是CIFAR-10原因很简单CIFAR-100的100个类别划分比CIFAR-10的10个大类要细得多同一个超大类下面还有20个细粒度子类比如鱼类下面分小型淡水鱼、大型淡水鱼等。这意味着模型不仅要学会区分这是鱼还要区分这是哪种鱼难度直接上了一个台阶。对于想深入理解图像分类、研究模型容量与泛化能力的人来说CIFAR-100是一个比CIFAR-10更合适的实验场。这个项目的定位很明确不是把某个SOTA论文复现到极致刷榜而是做一个多算法横向对比 工程化落地的分类框架。你可以在同一套数据管线、训练管线、评估管线里一键切换不同模型对比它们的精度、参数量、训练耗时。这份代码既适合刚入门的学生系统学习图像分类的完整流程也适合有经验的工程师快速验证某个新想法——把backbone替换掉就可以直接跑。1.1 CIFAR-100数据集深度解析CIFAR-100数据集共包含60000张32×32像素的彩色图像其中50000张用于训练10000张用于测试。100个类别每个类别恰好有600张图像训练集每类500张测试集每类100张。32×32的分辨率放到现在来看非常小但正因为小训练速度快、迭代试错成本低特别适合做算法对比和教学演示。需要特别注意的一点是CIFAR-100的类别组织是层次化的100个细粒度类别归属于20个超大类superclass比如鱼这个超大类下面有小型淡水鱼、大型淡水鱼等5个细分类别。torchvision.datasets.CIFAR100在加载数据时有一个参数superclass设置为True会返回20个超大类标签而不是100个细粒度标签。做多标签层次分类实验时这个参数会非常有用。数据集本身的图像质量参差不齐有些类别的图像即使人眼去看都非常模糊甚至无法准确判断类别这也意味着CIFAR-100的天花板并不是100%。从我实测的结果来看人的视觉准确率大概在90%左右这也解释了为什么很多模型在CIFAR-100上做到80%以上精度就已经是非常好的成绩了。1.2 为什么用PyTorch而不是其他框架选择PyTorch作为实现框架理由非常务实第一动态计算图让调试和原型验证极其顺畅。训练过程中想打印中间层的梯度、临时修改网络结构、用torch.autograd做梯度检查这些操作在PyTorch里都是顺手的事。我在开发过程中频繁需要检查某个层的输出shape直接在forward里加print就行这种开发体验是静态图框架很难比的。第二生态成熟度极高。torchvision提供了CIFAR-100数据集的直接下载接口和标准预处理管道模型的预训练权重也随手可得。社区里关于PyTorch的教程、踩坑记录非常丰富遇到问题搜索解决方案的效率非常高。第三代码可读性好适合教学和二次开发。PyTorch的编程范式非常Pythonic——用nn.Module定义网络结构、用torch.utils.data.DataLoader加载数据、用optimizer.zero_grad()手动清空梯度。每一步都清晰透明没有黑魔法学习者可以逐步跟踪数据流动的细节。2. 环境搭建与数据流水线构建2.1 PyTorch环境搭建的踩坑实录在开始编码之前先把环境搭好。我这次用的是Windows Anaconda的组合PyTorch版本是2.xCUDA版本为11.8。这里有几个值得分享的细节创建独立的conda环境是第一步。很多初学者习惯直接在base环境里pip install但不同项目的PyTorch版本、CUDA版本、Python版本很容易冲突。一个干净的环境是保障开发效率的基础。推荐做法conda create -n cifar100 python3.9 conda activate cifar100 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118注意PyTorch 2.x已经默认支持weights_only参数的调整在加载预训练权重时如果遇到FutureWarning: You are usingtorch.loadwithweights_onlyFalse的提示建议显式指定weights_onlyTrue或者直接用torchvision提供的model.load_state_dict(torch.load(...))方式加载避免安全隐患。验证CUDA是否可用这一步看似简单但非常重要import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))如果torch.cuda.is_available()返回False大概率是CUDA版本和PyTorch的cu版本不匹配或者NVIDIA驱动太旧。需要更新驱动到适配版本并重新安装对应CUDA版本的PyTorch。2.2 数据处理与增强策略CIFAR-100的图像是32×32的RGB图像原始数据量很小直接在训练时把所有图像读入内存完全可行。但数据增强策略直接影响模型的最终精度这是整个项目中性价比最高的调优手段之一。我采用的数据增强管线如下import torchvision.transforms as transforms train_transform transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean(0.5071, 0.4867, 0.4408), std(0.2675, 0.2565, 0.2761)) ]) test_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean(0.5071, 0.4867, 0.4408), std(0.2675, 0.2565, 0.2761)) ])这里有几个关键细节需要解释RandomCrop(padding4)先把32×32的图像padding到40×40再随机裁剪回32×32。这等于让模型看到了图像的平移不变性能有效缓解过拟合。CIFAR-10上这个操作能提升几个点的精度CIFAR-100上效果同样明显。Mean和Std的取值CIFAR-100和CIFAR-10的均值和标准差是不一样。很多初学者直接套用CIFAR-10的(0.4914, 0.4822, 0.4465)但实际上CIFAR-100的官方统计是(0.5071, 0.4867, 0.4408)和(0.2675, 0.2565, 0.2761)。用错数值会导致输入分布偏移虽然不会让模型完全无法训练但会让收敛变慢、精度损失1-2个点。要不要用CutOut/AutoAugment我实测过CutOut随机遮挡图像中的一个方形区域在CIFAR-100上大约能提升0.5-1个点的精度但对训练时长的增加也比较明显。AutoAugment效果更好但计算开销更大在入门项目中不建议一上来就上这些高级策略——先把基础增强用好再逐步加码。数据加载使用torch.utils.data.DataLoader设置num_workers4Windows上建议不要超过CPU核心数太高反而会报错batch_size128。pin_memoryTrue可以把数据从CPU内存直接搬运到GPU显存减少传输时间。2.3 标签映射与类别可视化CIFAR-100的类别名称存放在数据集的classes属性中包含100个字符串标签比如apple、beaver、telephone等。在训练之前最好先打印出类别列表了解你需要区分哪些东西。另外建议做一个类别分布可视化把每个类别的样本数画出来看看是否有不均衡问题。CIFAR-100是严格均衡的每类600张但如果后续你把项目扩展到其他数据集这一步不能省。3. 多种算法的实现与对比3.1 自定义CNN基线模型一个干净的基线模型非常重要它决定了你的地板在哪里。我从最经典的卷积全连接结构出发设计了一个轻量级CNNimport torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes100): super(SimpleCNN, self).__init__() self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(128, 256), nn.ReLU(inplaceTrue), nn.Dropout(0.5), nn.Linear(256, num_classes) ) def forward(self, x): return self.classifier(self.features(x))这个网络的参数量大约在60万左右在CIFAR-100上单卡训练40个epoch就可以达到大约55%-60%的Top-1准确率。它的作用不是刷精度而是作为后续所有算法的对照基准——所有改进模型都应该在这个基础上做增量对比。为什么用BatchNormBatchNorm在卷积层之后能显著加速收敛缓解梯度消失还能起到一定的正则化效果。没有BatchNorm的网络在CIFAR-100上训练会明显更慢而且对初始化权重更敏感。3.2 ResNet系列解决深层网络的训练瓶颈ResNet残差网络是图像分类任务上最经典的里程碑式工作其核心思想是引入残差学习——让网络学习输入和输出之间的残差映射F(x) H(x) - x而不是直接学习H(x)。这样可以有效缓解深层网络的退化问题即训练集上误差反而更高的现象。我实现了ResNet-18和ResNet-34两个版本核心就是残差块class BasicBlock(nn.Module): def __init__(self, in_channels, out_channels, stride1): super(BasicBlock, self).__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.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 x out F.relu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) out self.shortcut(identity) return F.relu(out)Shortcut连接为什么用1×1卷积当输入和输出的通道数不一致比如从64通道变成128通道或者特征图尺寸减半stride2时需要用1×1卷积做维度匹配这样才能保证相加操作合法。这里需要注意biasFalse——BatchNorm本身会做均值平移卷积层如果再带偏置就冗余了。ResNet-18在CIFAR-100上大约可以做到72%-75%的Top-1准确率ResNet-34能再提升1-2个点但训练时间也相应增加。相比之下ResNet-50使用了Bottleneck结构1×1降维→3×3卷积→1×1升维在更深的网络里计算效率更高但在CIFAR-100这种小尺寸图像上优势不明显反而容易过拟合。3.3 DenseNet特征复用的参数效率革命DenseNet的核心思想是稠密连接——每一层都与其前面所有层在通道维度上直接连接。和ResNet的求和相加不同DenseNet用的是拼接concatenation这样梯度可以直接从最后一层传到最前面的层有效缓解梯度消失。DenseNet块的核心实现class DenseBlock(nn.Module): def __init__(self, num_layers, in_channels, growth_rate): super(DenseBlock, self).__init__() self.layers nn.ModuleList() for i in range(num_layers): self.layers.append(BottleneckLayer( in_channels i * growth_rate, growth_rate )) def forward(self, x): for layer in self.layers: new_features layer(x) x torch.cat([x, new_features], dim1) return x class BottleneckLayer(nn.Module): def __init__(self, in_channels, growth_rate): super(BottleneckLayer, self).__init__() self.bn1 nn.BatchNorm2d(in_channels) self.conv1 nn.Conv2d(in_channels, 4 * growth_rate, 1, biasFalse) self.bn2 nn.BatchNorm2d(4 * growth_rate) self.conv2 nn.Conv2d(4 * growth_rate, growth_rate, 3, padding1, biasFalse) def forward(self, x): out F.relu(self.bn1(x)) out self.conv1(out) out F.relu(self.bn2(out)) out self.conv2(out) return outGrowth Rate增长率是最关键的超参数它决定了每一层新增多少个通道。常用值有12、24、32。growth_rate越大模型表达能力越强但参数量和计算量也越大。我用growth_rate12的DenseNet-BCBC代表Bottleneck Compression压缩因子为0.5参数量约0.8M精度可以到75%左右和ResNet-34相当但参数少了很多。DenseNet在CIFAR-100上的实际体验是收敛比ResNet慢但最终精度和稳定性都更好。特别是在训练后期DenseNet的loss下降曲线更加平稳不太容易反弹。3.4 训练配置优化器、损失函数与学习率策略训练配置对最终效果的影响不亚于模型结构本身。我这次对比了两种优化器方案SGD Momentum CosineAnnealingLR经典方案optimizer torch.optim.SGD(model.parameters(), lr0.1, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max200)AdamW OneCycleLR现代方案optimizer torch.optim.AdamW(model.parameters(), lr0.001, weight_decay0.05) scheduler torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr0.1, epochs200, steps_per_epochlen(train_loader) )实测下来的感受用SGD CosineAnnealing在CIFAR-100上的精度通常更高通用性也更强。AdamW收敛确实更快但最终精度往往比SGD略低1-2个点而且对weight_decay的取值非常敏感。OneCycleLR在训练初期用较低学习率热身中期上升到最大学习率后期再余弦衰减到接近0这种方式在CIFAR-100上能比固定学习率快将近一倍。设计一个热身体验如果你刚开始跑可以先用SGD学习率0.1、批大小128、训练40个epoch看趋势如果想快速验证一个模型能不能work就用AdamW OneCycleLR训练20个epoch就够判断了。4. 训练流程与监控4.1 训练循环标准实现完整的训练循环分三个部分训练、验证、测试。这是我整理的标准框架后面所有模型都复用了这套代码def train_one_epoch(model, train_loader, criterion, optimizer, device): model.train() running_loss 0.0 correct 0 total 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) _, predicted outputs.max(1) total labels.size(0) correct predicted.eq(labels).sum().item() epoch_loss running_loss / len(train_loader.dataset) epoch_acc 100.0 * correct / total return epoch_loss, epoch_acc def evaluate(model, test_loader, criterion, device): model.eval() running_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) running_loss loss.item() * images.size(0) _, predicted outputs.max(1) total labels.size(0) correct predicted.eq(labels).sum().item() test_loss running_loss / len(test_loader.dataset) test_acc 100.0 * correct / total return test_loss, test_acc一个关键细节训练模式下必须记得model.train()推理/评估时必须model.eval()。train()会启用Dropout和BatchNorm的训练行为使用当前批次的统计数据eval()会关闭Dropout并改用BatchNorm的全局统计。忘了切换的话推理结果会忽高忽低尤其是带了Dropout的模型。4.2 实验日志与可视化我习惯用tensorboard记录训练过程观察loss曲线和accuracy曲线来及时调整训练策略。启动方式很简单from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(log_dirruns/resnet34) # 在每个epoch结束后记录 writer.add_scalar(Loss/train, train_loss, epoch) writer.add_scalar(Loss/test, test_loss, epoch) writer.add_scalar(Accuracy/train, train_acc, epoch) writer.add_scalar(Accuracy/test, test_acc, epoch)判断模型是否健康的几个信号训练loss应该是平滑下降的没有剧烈震荡测试accuracy应该稳步上升到后期逐渐平缓。如果测试loss出现了U形曲线——先降后升说明已经过拟合了需要加强正则化或者提前停止训练。4.3 训练好的模型如何保存与加载保存模型的完整状态字典是最推荐的方式torch.save(model.state_dict(), checkpoints/resnet34_cifar100.pth)加载模型时需要先实例化一个结构完全相同的模型再加载权重model ResNet(BasicBlock, [2, 2, 2, 2], num_classes100) model.load_state_dict(torch.load(checkpoints/resnet34_cifar100.pth)) model model.to(device) model.eval()这里有一个容易踩的坑如果在保存模型之前调用了.to(device)权重会包含CUDA张量的信息在无GPU的环境加载时会报错。稳妥的方式是保存时用model.cpu().state_dict()或者加载时指定map_locationcpumodel.load_state_dict(torch.load(checkpoints/resnet34_cifar100.pth, map_locationtorch.device(cpu)))这样模型在任何环境都能成功加载。5. 实验对比与结果分析5.1 各模型精度-参数量-速度对比下面是我在统一配置下SGD优化器、CosineAnnealing、200个epoch、batch_size128跑出来的结果用NVIDIA RTX 3060 12G训练模型Top-1准确率参数量单epoch耗时备注SimpleCNN56.3%0.6M8s基线模型ResNet-1874.1%11.2M18s经典方案ResNet-3475.6%21.3M28s加深有效DenseNet-BC(12)75.2%1.0M25s参数效率最高DenseNet-BC(24)76.8%15.3M40s精度最高从结果可以看出几件事DenseNet-BC(12)是最具性价比的选择。它只有1M参数量精度却和ResNet-34差不多这说明特征复用机制在小数据集上非常有优势。但DenseNet的训练速度偏慢因为拼接操作带来的内存开销更大实测在3060上会明显感受到训练更卡。ResNet-18是主流基准。参数适中、精度不差、训练快适合做大多数实验的默认backbone。如果想在CIFAR-100上快速验证一个新想法ResNet-18是第一选择。5.2 训练曲线分析重点分析一下训练曲线背后的信息。ResNet-18的loss曲线在200个epoch中呈现三个典型阶段阶段一epoch 0-30模型从随机初始化开始快速学习训练loss从4.6约等于随机猜测100类的交叉熵损失快速下降到1.5左右测试精度从1%飙升到50%。这个阶段学习率尚未衰减模型在快速收敛。阶段二epoch 30-120训练loss继续下降但速度放缓测试精度从50%缓慢爬升到65%左右。这个阶段最容易出现的问题是过拟合——训练loss继续降但测试精度涨幅趋缓。解决方案是加强数据增强或增加Dropout比例。阶段三epoch 120-200CosineAnnealing学习率已经衰减到很低的水平模型在做精细化调整测试精度从65%提升到74%。这个阶段看似提升幅度不大但对每个epoch的学习率设置非常敏感学习率太高会导致loss反弹。5.3 错误案例分析混淆矩阵分析是理解模型短板的好方法。我随机抽取了1000张测试图像做了逐类别的错误分析发现几个规律细粒度类别最容易混淆比如枫树和橡树、苹果和梨这类外观高度相似的类别模型经常出错。这类错误本质上是信息量不足——32×32分辨率下树干纹理、叶子形状的细节已经完全丢失了人眼也不一定分得清。语义上的类别不对称现象模型倾向于把罕见特征预测为常见特征。比如公交车和卡车之间容易混淆但公交车误判为卡车的比例远高于反过来。这说明数据集中某些类别的特征分布更集中模型对这些类别更自信。小物体类别准确率偏低比如蜗牛、青蛙这类在32×32图像中特征区域很小的类别准确率明显低于汽车、飞机这类有大面积规则形状的类别。这个观察也解释了为什么目标检测领域对CIFAR这种小图数据集不感冒——分辨率本身就是硬伤。6. 训练中的常见问题与排查实录6.1 显存不足OOM问题训练CIFAR-100时显存不足的情况比想象中常见尤其是跑DenseNet或者batch_size设得比较大的时候。我的排查顺序第一步降低batch_size。把batch_size从128降到64如果显存不够继续降到32。这是最直接的解决方案但会影响BatchNorm的统计稳定性可能需要配合降低学习率。第二步检查模型是否一直在累积梯度图。如果代码里忘了optimizer.zero_grad()每一次backward都会叠加计算图显存占用会持续增长直到爆掉。在训练循环里确认这一步一定在loss.backward()之前optimizer.zero_grad() loss.backward() optimizer.step()第三步使用梯度累积模拟更大的batch。如果单卡显存太小可以用梯度累积技巧先把forward和backward跑几次每N次才做一次optimizer.step()等效于扩大了N倍的batch_sizeaccumulation_steps 4 for i, (images, labels) in enumerate(train_loader): outputs model(images) loss criterion(outputs, labels) loss loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()6.2 模型不收敛的排查清单如果训练了十几个epochloss还是纹丝不动或者直接变成NaN按下面的顺序逐项排查检查数据归一化。ToTensor()会把图像从[0,255]缩放到[0,1]但如果忘了做Normalize输入值过大或分布不均会导致梯度爆炸。检查transform流程是否完整。检查学习率是否过大或过小。学习率太大会导致loss变成NaN学习率太小会导致loss几乎不下降。一个通用的调试方法从学习率0.01开始每次阶梯性缩小10倍看哪个区间loss能正常下降。检查标签是否从0开始连续编码。CIFAR-100的标签是0-99的整数但如果你自定义数据集时标签从1开始或者中间有类别缺失nn.CrossEntropyLoss会报错或者计算出错误的loss。打印labels.min()和labels.max()确认一下。检查模型最后输出层维度是否等于类别数。100个类别最后的nn.Linear(last_channels, 100)如果写成128或者64CrossEntropyLoss的维度不匹配会直接报错。6.3 过拟合的应对策略CIFAR-100的训练集只有5万张图模型又动辄上百万参数过拟合几乎是必然的。我常用的正则化手段按优先级排列数据增强效果最好RandomCrop、HorizontalFlip、ColorJitter、CutOut这些操作相当于无限扩充了训练集。在CIFAR-100上完整的数据增强管线能带来10个点以上的精度提升。Weight DecaySGD的weight_decay设成5e-4是CIFAR系列任务的标准配置。weight_decay太大1e-2会导致模型欠拟合太小1e-5几乎不起作用。Label Smoothing把one-hot标签换成平滑分布比如真实类别概率为0.9其余类别均分0.1/99。这个技巧能防止模型对训练集过度自信提升泛化能力。criterion nn.CrossEntropyLoss(label_smoothing0.1)Early Stopping监控验证集的loss如果连续10个epoch没有改善就停止训练并回滚到最佳模型。这个机制在PyTorch里需要手动实现用列表记录历史最佳loss每次epoch结束判断是否更新。6.4 推理阶段的工程优化训练完了模型最终要落地上线。推理阶段也有很多细节批处理推理比单张推理快得多。即使只有100张图要预测也建议全部堆到一批里利用GPU的并行能力。如果必须在CPU环境推理可以用torch.set_num_threads(4)控制线程数或者用torch.jit.script(model)做图优化。导出为ONNX是跨平台部署的常见选择。如果之后要在手机上或者边缘设备上跑用ONNX Runtime比直接跑PyTorch轻量很多dummy_input torch.randn(1, 3, 32, 32).to(device) torch.onnx.export(model, dummy_input, resnet34_cifar100.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}})动态轴设置允许推理时随意指定batch大小这在服务端部署时非常实用。7. 扩展方向与个人心得7.1 从CIFAR-100到更大数据集的迁移CIFAR-100虽然只有32×32分辨率但其中学到的模型设计经验可以直接迁移到更高分辨率的数据集上。核心迁移点包括网络结构的适配原始ResNet-18的输入是3×32×32如果迁移到ImageNet3×224×224只需要调整网络第一层的stride和后续的池化参数。torchvision提供的ResNet18模型默认自带ImageNet适配版本可以直接加载预训练权重微调。数据增强管线的复用CIFAR-100的RandomCropHorizontalFlipNormalize这套组合在ImageNet等数据集上同样适用只需要相应调整crop的大小和padding值。训练策略的迁移SGDMomentumCosineAnnealing是通用性最强的训练配置从CIFAR-100到ImageNet一直适用。如果你想在更大数据集上做迁移学习直接在CIFAR-100上训练好的backbone权重作为初始化再解冻一部分层做微调能省下大量训练时间。7.2 踩过几次坑之后的一些经验总结做这个CIFAR-100多算法分类项目前前后后花了我大概两周时间。有几件事如果一开始就知道会少走很多弯路第一先定好对比实验的统一配置再动手。我一开始各个模型用了不同的训练超参最后对比结果根本没法公平地说明模型之间的差异。后来统一了epoch数、优化器、学习率策略、数据增强参数才得到一份能用的对比表。做算法对比实验控制变量永远是第一原则。第二不要迷信更深的网络。在CIFAR-100这种小数据集上ResNet-50和ResNet-18的精度差距可能就1个点左右但训练时间翻了一倍。先跑基线看清瓶颈在哪里再决定要不要加深网络。第三数据增强和调参带来的收益往往大于换模型。我在ResNet-18上加CutOut Label Smoothing精度从74.1%提升到76.5%这个提升幅度比从ResNet-18换成ResNet-3474.1%→75.6%还要大。所以如果你有个模型效果不理想先别急着换模型把增强和训练策略榨干再动结构。第四日志和可视化工具要尽早搭好。前期没有用TensorBoard每个epoch在终端里print一下loss和精度模型一多就乱了。后来统一用SummaryWriter记录所有实验的指标不同模型用不同颜色画到一张图上对比起来一目了然。实验管理这件事越早做越省心。7.3 这个项目后续还可以怎么扩展如果你想把CIFAR-100分类项目往深处做有几个方向值得尝试引入注意力机制SE模块Squeeze-and-Excitation是一个只需几行代码就能嵌入的通道注意力模块在ResNet的每个残差块里加上SE可以提升1个点左右的精度。更现代的Efficient Attention、CBAM也都是低成本的增强方案。尝试知识蒸馏用一个精度更高的大模型比如DenseNet-BC(24)76.8%当老师去教一个轻量学生模型比如SimpleCNN56.3%。好的蒸馏配置可以让学生模型精度提升到65%左右这在模型压缩场景下非常有价值。结合半监督学习CIFAR-100只有5万张有标签数据但我们可以用自监督预训练的方式让模型在无标签数据上先学习特征表示再微调。FixMatch这类半监督算法在CIFAR-100的设定下能有非常亮眼的表现。做分类项目的过程本质上就是一个不断逼近数据-模型-训练策略最优解的过程。CIFAR-100这个数据集足够小、足够快正好适合你去试各种想法、积累直觉。等你在CIFAR-100上把ResNet、DenseNet这些经典结构都跑熟了再去做目标检测、语义分割、自监督学习会发现很多思想都是相通的。本文还有配套的精品资源点击获取
返回列表