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, padding=4),
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', train=True, download=True, transform=transform_train)
trainloader = torch.utils.data.DataLoader(
trainset, batch_size=128, shuffle=True, num_workers=2)
testset = torchvision.datasets.CIFAR10(
root='./data', train=False, download=True, transform=transform_test)
testloader = torch.utils.data.DataLoader(
testset, batch_size=100, shuffle=False, num_workers=2)
数据增强是提升模型泛化能力的关键。我们使用了随机裁剪(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, stride=1):
super(BasicBlock, self).__init__()
self.conv1 = nn.Conv2d(
in_planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3,
stride=1, padding=1, bias=False)
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_size=1, stride=stride, bias=False),
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-18:
class ResNet(nn.Module):
def __init__(self, block, num_blocks, num_classes=10):
super(ResNet, self).__init__()
self.in_planes = 64
self.conv1 = nn.Conv2d(3, 64, kernel_size=3,
stride=1, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(64)
self.layer1 = self._make_layer(block, 64, num_blocks[0], stride=1)
self.layer2 = self._make_layer(block, 128, num_blocks[1], stride=2)
self.layer3 = self._make_layer(block, 256, num_blocks[2], stride=2)
self.layer4 = self._make_layer(block, 512, num_blocks[3], stride=2)
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(), lr=0.1, momentum=0.9, weight_decay=5e-4)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=200)
我们使用带动量的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(f'Epoch: {epoch} | Loss: {train_loss/(batch_idx+1):.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(f'Test Loss: {test_loss/(batch_idx+1):.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, padding=4),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.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_size=3, stride=1, padding=1, bias=False)
# 替换原始ResNet的:
# self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)
# self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
这样可以保留更多空间信息。同时可以减小模型宽度,将初始通道数从64降到32,减少计算量。
5.3 标签平滑正则化
标签平滑(Label Smoothing)可以缓解过拟合:
criterion = nn.CrossEntropyLoss(label_smoothing=0.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, label='Train')
plt.plot(test_losses, label='Test')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()
plt.subplot(1,2,2)
plt.plot(train_accs, label='Train')
plt.plot(test_accs, label='Test')
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, annot=True, fmt='d', cmap='Blues')
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。
更多推荐




所有评论(0)