图像分类全流程
本教程共 60 篇 · 第 35 篇 · 更新于 2026-08-17 · 约 3 分钟阅读
本节目标:把数据、模型、训练、评估串成一条标准流水线,学会带验证阶段的训练模板、按类别分析错误,以及保存模型并对新图片做预测。
第 21 章我们在 CIFAR-10 上跑通了第一个分类器。那是个「能跑就行」的骨架,这章把每个环节补细:数据怎么增强、训练时怎么盯着验证集、模型怎么存、新图怎么预测。图像分类(Image Classification)的任务就一句话:给一张图,从固定的类别表里挑一个。
流程总览
完整流程分六步:
- 准备数据:加载、预处理、数据增强。
- 定义模型:搭网络或拿预训练模型。
- 选损失函数与优化器。
- 训练:每个 epoch 里跑训练和验证两个阶段。
- 评估:总准确率、每类准确率、混淆矩阵。
- 保存模型,对新图片做预测。
第 1、2 步第 21 章讲过,这章快速带过,重点在 4、5、6 三步。
第一步:准备数据
继续用 CIFAR-10。关键区别是:训练集加随机增强,测试集只做确定性处理。
import torch
import torchvision
import torchvision.transforms as transforms
train_transform = transforms.Compose([
transforms.RandomHorizontalFlip(), # 随机水平翻转,增强
transforms.RandomCrop(32, padding=4), # 填充后随机裁剪,增强
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])
test_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=train_transform)
testset = torchvision.datasets.CIFAR10(
root='./data', train=False, download=True, transform=test_transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=64, shuffle=True)
testloader = torch.utils.data.DataLoader(testset, batch_size=64, shuffle=False)
classes = ('plane', 'car', 'bird', 'cat', 'deer',
'dog', 'frog', 'horse', 'ship', 'truck')
随机翻转、随机裁剪让每一轮看到的图都略有不同,等于变相扩充了数据,能防过拟合。测试集不增强,因为评估要的是稳定可比的结果。
第二步:定义模型
还是两层卷积的小网络,结构参考第 21 章:
import torch.nn as nn
import torch.nn.functional as F
class Net(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 6, 5)
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(6, 16, 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()
想省事也可以直接拿第 34 章的预训练模型换最后一层。两条路都能到终点,选哪条看你的数据量。
第三步:损失函数与优化器
多分类的标配:交叉熵损失加带动量的 SGD。
import torch.optim as optim
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)
学习率先给 0.001,loss 降不动再调。第 14、15 章讲过优化器和学习率调度的细节,这里不展开。不过建议顺手加一个调度器:每隔几个 epoch 把学习率缩小一次,后期收敛更稳。下一节的代码里就加上。
第四步:训练与验证
这一步是本章的重点:每个 epoch 里,先训练阶段更新参数,再验证阶段看泛化表现。验证时不更新参数,只看成绩:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
net.to(device)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5)
best_acc = 0.0
for epoch in range(5):
# ---- 训练阶段 ----
net.train()
running_loss = 0.0
for inputs, labels in trainloader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = net(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
# ---- 验证阶段(这里直接用测试集) ----
net.eval()
correct, total = 0, 0
with torch.no_grad():
for inputs, labels in testloader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = net(inputs)
_, predicted = torch.max(outputs, 1)
total += labels.size(0)
correct += (predicted == labels).sum().item()
acc = correct / total
print(f"Epoch {epoch+1}: train loss {running_loss/len(trainloader):.3f}, "
f"test acc {acc:.3f}")
if acc > best_acc:
best_acc = acc
scheduler.step() # 每个 epoch 结束时更新学习率
两个阶段各有一个开关:net.train() 打开 Dropout、BatchNorm 的训练行为,net.eval() 关闭它们;验证阶段还用 torch.no_grad() 关掉梯度,省显存又更快。acc > best_acc 时记下最好成绩,最后保存的就是它。
Note验证集和训练集必须是分开的两批数据。拿训练数据当验证,等于考试前偷看答案,分数再高也没意义。
第五步:按类别分析错误
总准确率是个平均数,会掩盖短板。哪几类容易认错?看每类准确率:
class_correct = [0. for _ in range(10)]
class_total = [0. for _ in range(10)]
net.eval()
with torch.no_grad():
for inputs, labels in testloader:
inputs = inputs.to(device)
outputs = net(inputs)
_, predicted = torch.max(outputs, 1)
predicted = predicted.cpu()
c = (predicted == labels) # 每个样本预测对没有
for i in range(len(labels)):
label = labels[i].item()
class_correct[label] += c[i].item()
class_total[label] += 1
for i in range(10):
print(f"{classes[i]:5s}: {100 * class_correct[i] / class_total[i]:.0f}%")
哪类分数特别低,就针对它加数据或加增强。想看得更细,用混淆矩阵(Confusion Matrix):横轴是预测类别,纵轴是真实类别。对角线是分对的,非对角线是互相认错的。矩阵里数字最大的非对角线格子,就是最容易被搞混的一对:
import numpy as np
preds, labels = [], [] # 收集阶段见上面循环,把每批 predicted 和 labels 追加进来
cm = np.zeros((10, 10), dtype=int)
for p, t in zip(preds, labels):
cm[t, p] += 1
print(cm)
画热力图可以用 matplotlib 的 imshow 或 seaborn 的 heatmap,颜色越深数字越大,看起来更直观。懒得自己算的话,scikit-learn 里有现成的:from sklearn.metrics import confusion_matrix, classification_report,把收集到的 preds、labels 传进去,连精确率、召回率、F1 一起算好。
第六步:保存模型与预测新图
训练完把 state_dict 存下来,预测时重建模型再加载,这是第 25 章的套路:
torch.save(net.state_dict(), './cifar_net.pth')
# 预测时:重建结构,再加载参数
net2 = Net()
net2.load_state_dict(torch.load('./cifar_net.pth'))
net2.eval()
对新图片做预测,注意两点:预处理要和训练时一致;输入要加一个 batch 维度。完整写法:
from PIL import Image
img = Image.open('./my_cat.png').convert('RGB')
img = test_transform(img) # 和测试集一样的处理
img = img.unsqueeze(0) # (1, 3, 32, 32),多出 batch 维
with torch.no_grad():
outputs = net2(img)
_, predicted = torch.max(outputs, 1)
print("预测结果:", classes[predicted.item()])
输出是 10 个类别的得分,torch.max 取最大那个的序号,再从 classes 里查名字。
整个流程在 CPU 上跑 5 个 epoch 大约需要几分钟到十几分钟,加数据增强会慢一些,换 GPU 就快得多。训练慢不用慌,先让它跑着,下一章的知识点可以先看起来。
Tip
torch.load默认只允许加载权重张量(weights_only=True),防的是恶意文件。自己保存的文件随便加载,来路不明的 .pth 文件要小心。
新手坑清单
流程走完,把容易踩的坑集中列一遍:
- 忘记
optimizer.zero_grad():梯度越积越多,训练直接发散。 - 验证阶段忘了
net.eval():Dropout 还在随机丢神经元,成绩偏低。 - 图片通道顺序搞反:PyTorch 要 C×H×W,PIL 图是 H×W×C。
ToTensor()会帮你转,自己处理时别弄反。 - 预测时忘了
unsqueeze(0):模型要 batch 维,不然报形状错误。 - 训练和预测的预处理不一致:训练归一化到 [-1,1],预测时也得一样,否则成绩稀烂。
小结
图像分类全流程六个步骤,核心是第四步那个「训练 + 验证」的模板:训练阶段更新参数,验证阶段只看成绩并记录最好模型。加上按类别分析错误、保存与预测,你就拥有了一套能直接套用到其他分类任务上的标准流程。这套骨架再往外扩:换数据、换模型、加调度器,万变不离其宗。下一个话题,我们来认识一个给 CNN 装「摆正」能力的模块——空间变换网络。