PyTorch 实战:训练 CIFAR-10 图像分类器
一、任务简介
图像分类是计算机视觉的入门任务。本文我们将使用 PyTorch 训练一个卷积神经网络(CNN),对 CIFAR-10 数据集中的 10 类彩色图像进行分类。你将学会:
- 下载和预处理 CIFAR-10 数据集
- 定义适用于彩色图像的 CNN 结构
- 训练模型并保存/加载模型
- 在测试集上评估模型性能(整体准确率和各类别准确率)
- 将训练过程迁移到 GPU 以加速
二、CIFAR-10 数据集介绍
CIFAR-10 是一个经典的图像分类数据集,包含 60000 张 32x32 的彩色图像(3 个通道),共 10 个类别:
| 类别 | 标签 |
|---|---|
| 0 | airplane(飞机) |
| 1 | automobile(汽车) |
| 2 | bird(鸟) |
| 3 | cat(猫) |
| 4 | deer(鹿) |
| 5 | dog(狗) |
| 6 | frog(青蛙) |
| 7 | horse(马) |
| 8 | ship(船) |
| 9 | truck(卡车) |
训练集 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 加速训练。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)