欢迎光临
我们一直在努力

初识深度学习——模型加载与推理

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

在前几篇博客中,我们从零构建了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 两种方式的对比

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

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

    模型加载完成后,还需要准备待推理的数据。代码中复用了训练时定义的 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) 取最大得分索引
    结果对比 预测值与真实值逐条比较
    赞(0)
    未经允许不得转载:171主机测评 » 初识深度学习——模型加载与推理
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址