一、引言:如何加载训练好的模型
在前几篇博客中,我们从零构建了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"))
步骤解析:
CNN():实例化模型,此时参数是随机初始化的。
torch.load("best2026-910.pth"):从文件中读取参数字典。
model.load_state_dict(…):将读取到的参数填入模型。
优点:
-
文件小,只存参数
-
灵活,可以加载到不同但结构相同的模型
-
是PyTorch推荐的标准做法
缺点:
-
必须知道模型结构,并正确定义
-
如果模型结构改变,旧参数可能无法加载
2.2 方式二:加载完整模型(TorchScript)
# 加载模型
model = torch.jit.load("best910.pth")
这是另一种加载方式。保存时使用 torch.jit.script(model) 和 torch.jit.save(),将模型结构、参数和计算图一起保存,加载时无需定义模型结构。
优点:
-
无需定义模型结构,直接加载即可用
-
可以跨平台部署
-
适合生产环境
缺点:
-
文件较大
-
某些动态结构可能无法脚本化
2.3 两种方式的对比
| 保存内容 | 仅参数 | 结构+参数+计算图 |
| 加载前提 | 需定义模型结构 | 无需定义 |
| 文件大小 | 小 | 大 |
| 部署灵活性 | 一般 | 高 |
| 推荐场景 | 研究、继续训练 | 生产、跨平台部署 |
三、推理前的准备:数据变换与数据集
模型加载完成后,还需要准备待推理的数据。代码中复用了训练时定义的 data_transforms 和 food_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_data 和 matplotlib:
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) 取最大得分索引 |
| 结果对比 | 预测值与真实值逐条比较 |


