首页 / PyTorch 入门教程 / 图像分类全流程

PyTorch 入门教程

图像分类全流程

本教程共 60 篇 · 第 35 篇 · 更新于 2026-08-17 · 约 3 分钟阅读

PyTorch图像分类训练流程混淆矩阵模型保存CIFAR-10

本节目标:把数据、模型、训练、评估串成一条标准流水线,学会带验证阶段的训练模板、按类别分析错误,以及保存模型并对新图片做预测。

第 21 章我们在 CIFAR-10 上跑通了第一个分类器。那是个「能跑就行」的骨架,这章把每个环节补细:数据怎么增强、训练时怎么盯着验证集、模型怎么存、新图怎么预测。图像分类(Image Classification)的任务就一句话:给一张图,从固定的类别表里挑一个。

流程总览

完整流程分六步:

  1. 准备数据:加载、预处理、数据增强。
  2. 定义模型:搭网络或拿预训练模型。
  3. 选损失函数与优化器。
  4. 训练:每个 epoch 里跑训练和验证两个阶段。
  5. 评估:总准确率、每类准确率、混淆矩阵。
  6. 保存模型,对新图片做预测。

第 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 文件要小心。

新手坑清单

流程走完,把容易踩的坑集中列一遍:

  1. 忘记 optimizer.zero_grad():梯度越积越多,训练直接发散。
  2. 验证阶段忘了 net.eval():Dropout 还在随机丢神经元,成绩偏低。
  3. 图片通道顺序搞反:PyTorch 要 C×H×W,PIL 图是 H×W×C。ToTensor() 会帮你转,自己处理时别弄反。
  4. 预测时忘了 unsqueeze(0):模型要 batch 维,不然报形状错误。
  5. 训练和预测的预处理不一致:训练归一化到 [-1,1],预测时也得一样,否则成绩稀烂。

小结

图像分类全流程六个步骤,核心是第四步那个「训练 + 验证」的模板:训练阶段更新参数,验证阶段只看成绩并记录最好模型。加上按类别分析错误、保存与预测,你就拥有了一套能直接套用到其他分类任务上的标准流程。这套骨架再往外扩:换数据、换模型、加调度器,万变不离其宗。下一个话题,我们来认识一个给 CNN 装「摆正」能力的模块——空间变换网络。