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 训练不收敛的可能原因

  1. 学习率设置不当 :尝试降低学习率(如从0.1降到0.01)或使用学习率查找器
  2. 数据预处理错误 :检查归一化参数是否正确,图像是否被正确处理
  3. 模型初始化问题 :确认没有错误的参数初始化导致梯度消失/爆炸
  4. 损失函数错误 :检查标签是否从0开始,与模型输出维度匹配

8.2 过拟合的解决方法

  1. 增加数据增强 :添加更多样的数据增强策略
  2. 使用更强的正则化 :增大权重衰减系数或添加Dropout层
  3. 早停(Early Stopping) :监控验证集性能,在开始下降时停止训练
  4. 模型简化 :减少模型层数或通道数

8.3 GPU内存不足的优化

  1. 减小批大小 :适当减小batch size,如从128降到64
  2. 使用梯度累积 :多次前向后累积梯度再更新参数
  3. 混合精度训练 :使用 torch.cuda.amp 自动混合精度
  4. 模型并行 :将模型拆分到多个GPU上

在实际项目中,我发现在CIFAR-10上ResNet-18的最佳batch size是128,太小会导致训练不稳定,太大又可能影响泛化性能。学习率初始设为0.1配合余弦退火通常能取得不错的效果。如果训练初期准确率一直不上升,可以检查数据加载是否正确,或者尝试先用小学习率(如0.01)预热几个epoch。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐