公司动态
PyTorch实战:基于ResNet-18的CIFAR-10图像分类
1. 项目概述在计算机视觉领域图像分类是最基础也最经典的任务之一。CIFAR-10数据集作为入门级的基准测试集包含了10个类别的6万张32x32像素彩色图像。这个项目展示了如何使用PyTorch框架基于ResNet-18架构实现一个完整的图像分类训练流程。我选择ResNet-18作为基础模型有几个考虑首先它的深度适中18层在CIFAR-10这样的小尺寸图像上既不会欠拟合也不会过拟合其次残差连接的设计让深层网络更容易训练最后作为经典模型它有大量可参考的实现和预训练权重。对于刚入门的开发者来说这个组合能快速验证想法并看到实际效果。2. 环境准备与数据加载2.1 基础环境配置推荐使用Python 3.8和PyTorch 1.10版本。以下是必需的依赖包pip install torch torchvision matplotlib tqdm如果使用GPU加速训练需要额外安装对应版本的CUDA工具包。可以通过nvidia-smi命令查看显卡支持的CUDA版本然后安装匹配的PyTorch版本。2.2 CIFAR-10数据集处理PyTorch的torchvision已经内置了CIFAR-10的下载和加载功能import torchvision import torchvision.transforms as transforms 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)), ]) trainset torchvision.datasets.CIFAR10( root./data, trainTrue, downloadTrue, transformtransform_train) trainloader torch.utils.data.DataLoader( trainset, batch_size128, shuffleTrue, num_workers2) testset torchvision.datasets.CIFAR10( root./data, trainFalse, downloadTrue, transformtransform_test) testloader torch.utils.data.DataLoader( testset, batch_size100, shuffleFalse, num_workers2)数据增强是提升模型泛化能力的关键。我们使用了随机裁剪RandomCrop和水平翻转RandomHorizontalFlip两种增强方式。归一化参数采用的是CIFAR-10数据集的全局均值和标准差。3. ResNet-18模型实现3.1 基础残差块设计ResNet的核心是残差连接Residual Connection下面是基础残差块的实现import torch.nn as nn import torch.nn.functional as F class BasicBlock(nn.Module): expansion 1 def __init__(self, in_planes, planes, stride1): super(BasicBlock, self).__init__() self.conv1 nn.Conv2d( in_planes, planes, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(planes) self.conv2 nn.Conv2d(planes, planes, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(planes) self.shortcut nn.Sequential() if stride ! 1 or in_planes ! self.expansion*planes: self.shortcut nn.Sequential( nn.Conv2d(in_planes, self.expansion*planes, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(self.expansion*planes) ) def forward(self, x): out F.relu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) out self.shortcut(x) out F.relu(out) return out每个残差块包含两个3x3卷积层中间有批归一化BatchNorm和ReLU激活。shortcut连接处理了输入输出维度不匹配的情况通过1x1卷积调整维度。3.2 完整ResNet-18架构基于上述基础块我们可以构建完整的ResNet-18class ResNet(nn.Module): def __init__(self, block, num_blocks, num_classes10): super(ResNet, self).__init__() self.in_planes 64 self.conv1 nn.Conv2d(3, 64, kernel_size3, stride1, padding1, biasFalse) self.bn1 nn.BatchNorm2d(64) self.layer1 self._make_layer(block, 64, num_blocks[0], stride1) self.layer2 self._make_layer(block, 128, num_blocks[1], stride2) self.layer3 self._make_layer(block, 256, num_blocks[2], stride2) self.layer4 self._make_layer(block, 512, num_blocks[3], stride2) self.linear nn.Linear(512*block.expansion, num_classes) def _make_layer(self, block, planes, num_blocks, stride): strides [stride] [1]*(num_blocks-1) layers [] for stride in strides: layers.append(block(self.in_planes, planes, stride)) self.in_planes planes * block.expansion return nn.Sequential(*layers) def forward(self, x): out F.relu(self.bn1(self.conv1(x))) out self.layer1(out) out self.layer2(out) out self.layer3(out) out self.layer4(out) out F.avg_pool2d(out, 4) out out.view(out.size(0), -1) out self.linear(out) return out def ResNet18(): return ResNet(BasicBlock, [2,2,2,2])与原始ResNet论文相比这里做了两处适配CIFAR-10的修改1) 去掉了第一个7x7卷积和最大池化层直接使用3x3卷积2) 最后的平均池化大小改为4因为CIFAR-10经过多次下采样后特征图尺寸为4x4。4. 训练流程实现4.1 训练超参数设置import torch.optim as optim device cuda if torch.cuda.is_available() else cpu net ResNet18().to(device) criterion nn.CrossEntropyLoss() optimizer optim.SGD(net.parameters(), lr0.1, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max200)我们使用带动量的SGD优化器初始学习率设为0.1配合余弦退火学习率调度器。权重衰减L2正则化设为5e-4防止过拟合。损失函数使用交叉熵损失这是多分类问题的标准选择。4.2 训练循环实现完整的训练循环包括前向传播、损失计算、反向传播和参数更新def train(epoch): net.train() train_loss 0 correct 0 total 0 for batch_idx, (inputs, targets) in enumerate(trainloader): inputs, targets inputs.to(device), targets.to(device) optimizer.zero_grad() outputs net(inputs) loss criterion(outputs, targets) loss.backward() optimizer.step() train_loss loss.item() _, predicted outputs.max(1) total targets.size(0) correct predicted.eq(targets).sum().item() acc 100.*correct/total print(fEpoch: {epoch} | Loss: {train_loss/(batch_idx1):.3f} | Acc: {acc:.3f}%) def test(epoch): net.eval() test_loss 0 correct 0 total 0 with torch.no_grad(): for batch_idx, (inputs, targets) in enumerate(testloader): inputs, targets inputs.to(device), targets.to(device) outputs net(inputs) loss criterion(outputs, targets) test_loss loss.item() _, predicted outputs.max(1) total targets.size(0) correct predicted.eq(targets).sum().item() acc 100.*correct/total print(fTest Loss: {test_loss/(batch_idx1):.3f} | Acc: {acc:.3f}%) return acc每个epoch结束后我们会在测试集上评估模型性能。注意训练和测试时要分别调用net.train()和net.eval()这会影响到BatchNorm和Dropout等层的行为。4.3 学习率调度策略我们使用余弦退火学习率调度这是近年来图像分类任务中的常用策略for epoch in range(200): train(epoch) test_acc test(epoch) scheduler.step() # 保存最佳模型 if test_acc best_acc: best_acc test_acc torch.save(net.state_dict(), best_model.pth)余弦退火让学习率从初始值缓慢下降到0模拟了重启的效果有助于跳出局部最优。通常训练200个epoch就能达到不错的性能。5. 模型优化与调参技巧5.1 数据增强策略优化除了基础的随机裁剪和水平翻转还可以尝试以下增强策略transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ])颜色抖动ColorJitter和随机旋转RandomRotation可以进一步提升模型鲁棒性。但要注意增强强度不宜过大否则会引入太多噪声。5.2 模型结构调整对于CIFAR-10这样的小尺寸图像可以调整ResNet的初始层self.conv1 nn.Conv2d(3, 64, kernel_size3, stride1, padding1, biasFalse) # 替换原始ResNet的 # self.conv1 nn.Conv2d(3, 64, kernel_size7, stride2, padding3, biasFalse) # self.maxpool nn.MaxPool2d(kernel_size3, stride2, padding1)这样可以保留更多空间信息。同时可以减小模型宽度将初始通道数从64降到32减少计算量。5.3 标签平滑正则化标签平滑Label Smoothing可以缓解过拟合criterion nn.CrossEntropyLoss(label_smoothing0.1)这会让真实标签从1变为0.9其他类别从0变为0.1/(num_classes-1)防止模型对训练标签过度自信。6. 结果分析与可视化6.1 训练曲线可视化使用Matplotlib绘制训练过程中的损失和准确率曲线import matplotlib.pyplot as plt plt.figure(figsize(12,4)) plt.subplot(1,2,1) plt.plot(train_losses, labelTrain) plt.plot(test_losses, labelTest) plt.xlabel(Epoch) plt.ylabel(Loss) plt.legend() plt.subplot(1,2,2) plt.plot(train_accs, labelTrain) plt.plot(test_accs, labelTest) plt.xlabel(Epoch) plt.ylabel(Accuracy) plt.legend() plt.show()典型的训练曲线应该显示训练损失稳步下降测试准确率逐步提升。如果出现训练准确率高但测试准确率低说明可能过拟合了。6.2 混淆矩阵分析混淆矩阵能直观展示模型在各类别上的表现from sklearn.metrics import confusion_matrix import seaborn as sns conf_mat confusion_matrix(all_targets, all_preds) plt.figure(figsize(10,8)) sns.heatmap(conf_mat, annotTrue, fmtd, cmapBlues) plt.xlabel(Predicted) plt.ylabel(Actual) plt.show()CIFAR-10中猫和狗、鹿和马等相似类别容易混淆。针对这些类别可以收集更多样本或设计特定的数据增强。7. 模型部署与应用7.1 模型保存与加载训练完成后保存整个模型或仅保存参数# 保存整个模型 torch.save(net, full_model.pth) # 仅保存参数推荐 torch.save(net.state_dict(), model_params.pth) # 加载模型 net ResNet18() net.load_state_dict(torch.load(model_params.pth))仅保存参数更灵活可以在加载时修改模型结构。完整的模型保存包含了类定义可能在不同环境中不兼容。7.2 单张图像预测实现一个简单的预测函数from PIL import Image def predict(image_path): img Image.open(image_path) img transform_test(img).unsqueeze(0).to(device) with torch.no_grad(): output net(img) _, predicted torch.max(output.data, 1) return classes[predicted.item()]注意输入图像需要经过与训练时相同的预处理流程。对于实际应用可以添加图像resize和中心裁剪等步骤。8. 常见问题与解决方案8.1 训练不收敛的可能原因学习率设置不当尝试降低学习率如从0.1降到0.01或使用学习率查找器数据预处理错误检查归一化参数是否正确图像是否被正确处理模型初始化问题确认没有错误的参数初始化导致梯度消失/爆炸损失函数错误检查标签是否从0开始与模型输出维度匹配8.2 过拟合的解决方法增加数据增强添加更多样的数据增强策略使用更强的正则化增大权重衰减系数或添加Dropout层早停Early Stopping监控验证集性能在开始下降时停止训练模型简化减少模型层数或通道数8.3 GPU内存不足的优化减小批大小适当减小batch size如从128降到64使用梯度累积多次前向后累积梯度再更新参数混合精度训练使用torch.cuda.amp自动混合精度模型并行将模型拆分到多个GPU上在实际项目中我发现在CIFAR-10上ResNet-18的最佳batch size是128太小会导致训练不稳定太大又可能影响泛化性能。学习率初始设为0.1配合余弦退火通常能取得不错的效果。如果训练初期准确率一直不上升可以检查数据加载是否正确或者尝试先用小学习率如0.01预热几个epoch。