一、引言:如何加载训练好的模型

在前几篇博客中,我们从零构建了CNN模型,用数据增强提升了泛化能力,并学会了在训练过程中保存最优模型。现在,我们手里已经有了两个模型文件:

best2026-910.pth:保存的模型参数(state_dict)

best910.pth:保存的完整TorchScript模型

但问题来了:训练好的模型,怎么拿来用? 总不能在每次预测时都重新训练一遍吧?

答案就是——加载模型,进行推理(Inference)。推理是指用训练好的模型对新的数据进行预测。本篇博客将基于一份完整的推理代码,讲解如何加载模型、如何准备数据、如何得到预测结果,并对比预测值与真实值,评估模型在测试集上的表现。

二、模型加载的两种方式

PyTorch提供了两种保存模型的方式,对应两种加载方式。代码中同时展示了这两种方法。

2.1 方式一:加载模型参数(state_dict)

这是PyTorch官方推荐的方式。保存时只保存了模型的参数(权重w和偏置b),加载时需要先定义模型结构,再加载参数。

# 定义模型结构(必须与训练时完全一致)
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
model = CNN().to(device)

# 加载参数
model.load_state_dict(torch.load("best2026-910.pth"))

步骤解析

  1. CNN():实例化模型,此时参数是随机初始化的。

  2. torch.load("best2026-910.pth"):从文件中读取参数字典。

  3. model.load_state_dict(...):将读取到的参数填入模型。

优点

  • 文件小,只存参数

  • 灵活,可以加载到不同但结构相同的模型

  • 是PyTorch推荐的标准做法

缺点

  • 必须知道模型结构,并正确定义

  • 如果模型结构改变,旧参数可能无法加载

2.2 方式二:加载完整模型(TorchScript)

# 加载模型
model = torch.jit.load("best910.pth")

这是另一种加载方式。保存时使用 torch.jit.script(model)torch.jit.save(),将模型结构、参数和计算图一起保存,加载时无需定义模型结构。

优点

  • 无需定义模型结构,直接加载即可用

  • 可以跨平台部署

  • 适合生产环境

缺点

  • 文件较大

  • 某些动态结构可能无法脚本化

2.3 两种方式的对比

对比项state_dictTorchScript
保存内容仅参数结构+参数+计算图
加载前提需定义模型结构无需定义
文件大小
部署灵活性一般
推荐场景研究、继续训练生产、跨平台部署

三、推理前的准备:数据变换与数据集

模型加载完成后,还需要准备待推理的数据。代码中复用了训练时定义的 data_transformsfood_dataset 类。

3.1 验证集变换

data_transforms = {
    'trainda': transforms.Compose([...]),   # 训练变换(含数据增强)
    'valid': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], 
                             std=[0.229, 0.224, 0.225])
    ]),
}

关键点:推理时使用的是 'valid' 变换,不能使用 'trainda'。因为训练变换中包含随机旋转、翻转、颜色抖动等数据增强操作,这些操作会引入随机性,导致同一张图片每次预测结果可能不同。推理时需要的是确定性的预处理。

3.2 自定义数据集类

food_dataset 负责读取 test.txt 中的图片路径和标签,并应用变换:

class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        # 读取文件,保存路径和标签
        ...
    def __len__(self):
        return len(self.imgs)
    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx])
        if self.transform:
            image = self.transform(image)
        label = torch.from_numpy(np.array(self.labels[idx], dtype=np.int64))
        return image, label

3.3 创建测试数据加载器

test_data = food_dataset(file_path='./test.txt', 
                         transform=data_transforms['valid'])
test_loader = DataLoader(test_data, batch_size=1, shuffle=True)

注意

  • batch_size=1:每次只处理一张图片,便于逐条记录预测结果。

  • shuffle=True:打乱顺序,但因为我们同时保存预测值和真实值,顺序不影响最终评估。

  • 如果只想快速评估准确率,可以设置更大的 batch_size 以加速。

四、模型推理:从输入到预测

加载模型和准备好数据后,就可以进行推理了。 test_true 函数完成了核心工作:

results = []    # 保存预测结果
labels = []     # 保存真实标签

def test_true(dataloader, model):
    with torch.no_grad():
        for x, y in dataloader:
            x, y = x.to(device), y.to(device)
            pred = model.forward(x)
            results.append(pred.argmax(1).item())
            labels.append(y.item())

test_true(test_loader, model)
print("预测值:\t", results)
print("真实值:\t", labels)

4.1 torch.no_grad()——关闭梯度计算

推理时不需要反向传播,因此可以关闭梯度计算:

with torch.no_grad():
    ...

作用

  • 减少内存消耗(不保存计算图)

  • 加快计算速度

  • 防止参数被意外修改

4.2 前向传播

pred = model.forward(x)

model.forward(x) 也可以简写为 model(x),PyTorch会自动调用 forward 方法。输出 pred 的形状为 (batch_size, 20),表示每张图片属于20个类别的得分。

4.3 获取预测类别

results.append(pred.argmax(1).item())
  • pred.argmax(1):在维度1(类别维度)上取最大值的索引,即预测的类别。

  • .item():将张量转为Python标量。

4.4 保存真实标签

labels.append(y.item())

y 是当前批次的真实标签张量,.item() 将其转为整数。

五、预测结果分析

运行后,会打印出两个列表:

预测值:	 [14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14]
真实值:	 [6, 12, 9, 17, 10, 4, 6, 11, 1, 15, 8, 7, 2, 16, 3, 16, 3, 13, 14, 14, 19, 16, 5, 10, 11, 5, 13, 17, 7, 18, 1, 18, 9, 19, 2, 8, 0, 3, 4]

通过对比这两个列表,我们可以:

5.1 计算准确率

correct = sum(1 for p, t in zip(results, labels) if p == t)
accuracy = correct / len(labels) * 100
print(f"准确率: {accuracy:.2f}%")

5.2 找出预测错误的样本

for i, (p, t) in enumerate(zip(results, labels)):
    if p != t:
        print(f"样本 {i}: 预测={p}, 真实={t}")

5.3 可视化预测结果

如果想查看具体图片,可以结合 test_datamatplotlib

import matplotlib.pyplot as plt

# 显示前9张图片及其预测结果
fig = plt.figure(figsize=(10, 10))
for i in range(9):
    img, true_label = test_data[i]
    pred_label = results[i]
    ax = fig.add_subplot(3, 3, i+1)
    ax.set_title(f"预测: {pred_label}, 真实: {true_label}")
    ax.axis('off')
    # 反标准化后显示
    img = img.permute(1, 2, 0).numpy()
    img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406]
    ax.imshow(img)
plt.show()

六、完整推理流程总结

import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import transforms
from PIL import Image
import numpy as np

# 1. 选择设备
device = "cuda" if torch.cuda.is_available() else "cpu"

# 2. 定义模型结构
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Sequential(
            nn.Conv2d(3, 16, 5, 1, 2),
            nn.ReLU(),
            nn.MaxPool2d(2),
        )
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2),
            nn.ReLU(),
            nn.Conv2d(32, 64, 5, 1, 2),
            nn.ReLU(),
            nn.MaxPool2d(2),
        )
        self.conv3 = nn.Sequential(
            nn.Conv2d(64, 128, 5, 1, 2),
            nn.ReLU(),
        )
        self.out = nn.Linear(128*64*64, 20)

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = x.view(x.size(0), -1)
        return self.out(x)

# 3. 加载模型参数
model = CNN().to(device)
model.load_state_dict(torch.load("best2026-910.pth"))
model.eval()

# 4. 数据准备
data_transforms = transforms.Compose([
    transforms.Resize([256, 256]),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], 
                         std=[0.229, 0.224, 0.225])
])

class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.imgs = []
        self.labels = []
        self.transform = transform
        with open(file_path) as f:
            for line in f:
                img_path, label = line.strip().split(' ')
                self.imgs.append(img_path)
                self.labels.append(label)

    def __len__(self):
        return len(self.imgs)

    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx])
        if self.transform:
            image = self.transform(image)
        label = torch.from_numpy(np.array(self.labels[idx], dtype=np.int64))
        return image, label

test_data = food_dataset('./test.txt', transform=data_transforms)
test_loader = DataLoader(test_data, batch_size=1, shuffle=True)

# 5. 推理
results = []
labels = []
with torch.no_grad():
    for x, y in test_loader:
        x, y = x.to(device), y.to(device)
        pred = model(x)
        results.append(pred.argmax(1).item())
        labels.append(y.item())

# 6. 输出结果
print("预测值:", results)
print("真实值:", labels)

七、总结

本篇博客围绕“模型加载与推理”这一主题,系统讲解了:

知识点核心内容
state_dict加载先定义模型结构,再加载参数
TorchScript加载直接加载完整模型,无需定义结构
模型评估模式model.eval() 固定参数
推理上下文torch.no_grad() 关闭梯度计算
预测类别pred.argmax(1) 取最大得分索引
结果对比预测值与真实值逐条比较

更多推荐