一、任务简介

图像分类是计算机视觉的入门任务。本文我们将使用 PyTorch 训练一个卷积神经网络(CNN),对 CIFAR-10 数据集中的 10 类彩色图像进行分类。你将学会:

  • 下载和预处理 CIFAR-10 数据集
  • 定义适用于彩色图像的 CNN 结构
  • 训练模型并保存/加载模型
  • 在测试集上评估模型性能(整体准确率和各类别准确率)
  • 将训练过程迁移到 GPU 以加速

二、CIFAR-10 数据集介绍

CIFAR-10 是一个经典的图像分类数据集,包含 60000 张 32x32 的彩色图像(3 个通道),共 10 个类别:

类别标签
0airplane(飞机)
1automobile(汽车)
2bird(鸟)
3cat(猫)
4deer(鹿)
5dog(狗)
6frog(青蛙)
7horse(马)
8ship(船)
9truck(卡车)

训练集 50000 张,测试集 10000 张。

三、数据下载与预处理

使用 torchvision 可以轻松下载并转换数据。我们将图像转为张量,并归一化到 [-1, 1] 范围。

import torch
import torchvision
import torchvision.transforms as transforms

# 定义预处理:转换为张量 + 归一化(均值0.5,标准差0.5)
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# 下载训练集
trainset = torchvision.datasets.CIFAR10(root='./data', train=True,
                                        download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=4,
                                          shuffle=True, num_workers=2)

# 下载测试集
testset = torchvision.datasets.CIFAR10(root='./data', train=False,
                                       download=True, transform=transform)
testloader = torch.utils.data.DataLoader(testset, batch_size=4,
                                         shuffle=False, num_workers=2)

# 类别名称
classes = ('plane', 'car', 'bird', 'cat', 'deer',
           'dog', 'frog', 'horse', 'ship', 'truck')

Windows 用户提示:如果出现 BrokenPipeError,请将 num_workers 设为 0。

四、可视化部分训练图像

为了直观感受数据,我们展示一个 batch 的图片:

import matplotlib.pyplot as plt
import numpy as np

def imshow(img):
    img = img / 2 + 0.5     # 反归一化到 [0,1]
    npimg = img.numpy()
    plt.imshow(np.transpose(npimg, (1, 2, 0)))
    plt.show()

# 获取一个 batch
dataiter = iter(trainloader)
images, labels = next(dataiter)

# 显示图片网格
imshow(torchvision.utils.make_grid(images))
# 打印标签
print(' '.join(f'{classes[labels[j]]:5s}' for j in range(4)))

五、定义卷积神经网络

由于 CIFAR-10 是彩色 3 通道图像,我们将第一层卷积的输入通道改为 3。网络结构:两个卷积层(5x5 卷积核) + 三个全连接层。

import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)      # 输入3通道,输出6通道
        self.pool = nn.MaxPool2d(2, 2)       # 池化窗口2x2
        self.conv2 = nn.Conv2d(6, 16, 5)     # 输入6,输出16
        # 经过两次卷积+池化后,特征图尺寸为 16 * 5 * 5(计算过程见下文)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 16 * 5 * 5)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

net = Net()
print(net)

尺寸计算:输入 32x32,第一个卷积(5x5,无 padding)输出 28x28,池化后 14x14;第二个卷积输出 10x10,池化后 5x5。因此展平维度为 1655=400。

六、定义损失函数和优化器

分类任务使用交叉熵损失(CrossEntropyLoss),优化器选择带动量的 SGD。

import torch.optim as optim

criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)

七、训练模型

训练 2 个 epoch(可根据需要增加),每个 epoch 遍历整个训练集。

for epoch in range(2):
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):
        inputs, labels = data
        # 梯度清零
        optimizer.zero_grad()
        # 前向传播
        outputs = net(inputs)
        # 计算损失
        loss = criterion(outputs, labels)
        # 反向传播
        loss.backward()
        # 更新参数
        optimizer.step()

        running_loss += loss.item()
        if (i+1) % 2000 == 0:
            print(f'[Epoch {epoch+1}, Batch {i+1:5d}] loss: {running_loss/2000:.3f}')
            running_loss = 0.0

print('训练完成!')

输出示例:

[Epoch 1, Batch  2000] loss: 2.227
[Epoch 1, Batch  4000] loss: 1.884
...
[Epoch 2, Batch 12000] loss: 1.291
训练完成!

八、保存模型

训练完成后,通常只保存模型的参数(state_dict)而不是整个对象。

PATH = './cifar_net.pth'
torch.save(net.state_dict(), PATH)

九、测试模型

1. 在单 batch 上测试

# 加载模型
net = Net()
net.load_state_dict(torch.load(PATH))

# 获取测试集的一个 batch
dataiter = iter(testloader)
images, labels = next(dataiter)

# 预测
outputs = net(images)
_, predicted = torch.max(outputs, 1)

# 显示真实标签和预测标签
print('GroundTruth: ', ' '.join(f'{classes[labels[j]]:5s}' for j in range(4)))
print('Predicted:   ', ' '.join(f'{classes[predicted[j]]:5s}' for j in range(4)))

2. 在整个测试集上评估准确率

correct = 0
total = 0
with torch.no_grad():   # 不计算梯度,节省内存
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

print(f'在 10000 张测试集上的准确率: {100 * correct / total:.1f}%')

通常可以达到 53% 左右的准确率(随机猜测仅 10%),说明模型确实学到了特征。

3. 分类别统计准确率

class_correct = [0.0] * 10
class_total = [0.0] * 10
with torch.no_grad():
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs, 1)
        c = (predicted == labels).squeeze()
        for i in range(4):   # batch_size=4
            label = labels[i]
            class_correct[label] += c[i].item()
            class_total[label] += 1

for i in range(10):
    print(f'{classes[i]:10s} 准确率: {100 * class_correct[i] / class_total[i]:.1f}%')

输出示例:

plane      准确率: 62.0%
car        准确率: 62.0%
bird       准确率: 45.0%
cat        准确率: 36.0%
deer       准确率: 52.0%
dog        准确率: 25.0%
frog       准确率: 69.0%
horse      准确率: 60.0%
ship       准确率: 70.0%
truck      准确率: 48.0%

可以看到模型对猫、狗的分类效果较差,对船、青蛙的分类效果较好。

十、在 GPU 上训练(加速)

如果你的电脑有 NVIDIA 显卡并已安装 CUDA,可以轻松将训练迁移到 GPU。

device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
print("使用设备:", device)

# 将网络转移到 GPU
net.to(device)

# 在训练循环中,将输入和标签也转移到 GPU
for epoch in range(2):
    for i, data in enumerate(trainloader, 0):
        inputs, labels = data
        inputs, labels = inputs.to(device), labels.to(device)
        # 其余代码不变 ...

只需添加 .to(device),PyTorch 就会自动在 GPU 上执行运算,训练速度显著提升。

十一、总结

  • 使用 torchvision 下载并预处理 CIFAR-10 数据集;
  • 为彩色图像设计 CNN 结构;
  • 完成完整的训练、保存、加载和测试流程;
  • 分析模型在不同类别上的表现;
  • 将模型迁移到 GPU 加速训练。
Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐