第10章:模型保存、加载与推理

🎉摘要:学习PyTorch模型保存与加载的最佳实践,掌握断点续训和独立推理脚本编写,实现MNIST手写数字识别模型的完整闭环,从训练到部署一步到位。

这一章的目标

  • 学会保存训练好的模型,下次不用重新训练

  • 学会加载已保存的模型,直接用它做预测

  • 支持断点续训——训练中断后可以从上次停下的地方继续

  • 写一个独立推理脚本,对任意手写图片做出预测

这一章学完,你就拥有了一个真正 "能用" 的 AI 模型。

保存模型

PyTorch 保存模型有两种方式,各有用处:

方式一:只保存参数(推荐)

# ====================== 模型权重保存 ======================
# torch.save:仅保存模型的网络权重参数 state_dict(推荐方式)
# state_dict 是字典结构,里面只存放卷积层、全连接层的权重、偏置等可训练参数;
# 不保存网络结构、优化器状态、epoch、损失等无关数据,文件体积小、兼容性好、轻量化部署
# 保存路径:当前目录下 mnist_cnn.pth
torch.save(model.state_dict(), "mnist_cnn.pth")

# ====================== 模型权重加载 ======================
# 1. 必须手动先实例化和训练时完全一致的网络结构 MNIST_CNN
# 单纯的权重文件里没有网络结构定义,需要先搭好一模一样的网络骨架
model = MNIST_CNN()

# 2. 从本地pth文件读取权重字典,并赋值给当前网络对应的层参数
# 加载后网络就拥有训练完成的最优权重,具备推理能力
model.load_state_dict(torch.load("mnist_cnn.pth"))

# 3. 切换模型为评估推理模式
# 关闭梯度计算相关逻辑,推理时不会自动计算梯度,节省显存、加快前向推理速度
model.eval()

优点:文件小、跨平台兼容性好。这是最推荐的方式

方式二:保存整个模型(不推荐)

# ====================== 模型保存 ======================
# 保存完整模型对象(网络结构 + 全部权重参数)
# 原理:直接序列化整个 model 实例,文件内部同时包含:
# 1)MNIST_CNN 类对应的完整网络结构定义;
# 2)卷积层、全连接层所有训练好的权重、偏置参数;
# 优点:加载时不需要提前手写搭建一模一样的网络结构,一行加载即可直接使用;
# 缺点:
# ① 文件体积比只存 state_dict 更大;
# ② 兼容性差,环境 PyTorch 版本、代码类定义发生改动时极易加载失败;
# ③ 工业部署标准场景一般不推荐该方式,仅适合本地临时快速调试
torch.save(model, "mnist_cnn_full.pth")

# ====================== 模型加载 ======================
# 读取之前通过 torch.save(model) 方式持久化保存的完整模型文件
# 文件内部封装了完整网络类结构 + 训练完毕的权重、偏置等参数,无需提前手动定义、实例化 MNIST_CNN 网络
# 加载逻辑:反序列化还原出训练时的整个模型实例,加载完成后模型已经自带训练好的参数,可以直接投入使用
# 缺陷:对代码环境、PyTorch版本、模型类定义依赖性极强,环境变动后大概率加载报错,正式工程不推荐该加载方案
model = torch.load("mnist_cnn_full.pth")

缺点:文件大、如果网络类定义改了可能加载不了。初学阶段别用这种方式。

📌 .pth 和 .pt 后缀

都是 PyTorch 模型文件的常用后缀,没有本质区别。.pth更常见,.pt是官方推荐。你用哪个都行。

保存最佳模型

训练过程中,我们通常保存 "在测试集上表现最好的那个模型",而不是 "最后一个":

# 初始化变量,用来记录训练至今为止测试集上拿到的最高准确率
best_acc = 0.0

# 按轮次循环执行完整训练流程
for epoch in range(num_epochs):
    # 执行一轮训练,获取本轮训练集平均损失、训练集准确率
    train_loss, train_acc = train_one_epoch(...)
    # 在测试集上做验证,得到本轮测试损失与测试准确率
    test_loss, test_acc = evaluate(...)

    # 判断:当前这一轮测试准确率超过历史最优记录
    if test_acc > best_acc:
        # 更新全局最优准确率为当前最新最优值
        best_acc = test_acc
        # 将当前模型权重保存为最优模型,只保存参数,占用体积小、工程常用
        torch.save(model.state_dict(), "best_model.pth")
        # 控制台打印提示,告知本次更新保存了最优权重以及对应的准确率
        print(f"  → 保存最佳模型!测试准确率: {test_acc:.2f}%")

断点续训

训练了几小时突然断电/崩溃了,重新从零开始训练太亏了。所以我们需要断点续训:保存的不仅是模型参数,还有优化器状态和 epoch 数。

保存 Checkpoint

# 构建训练断点续训所需的完整检查点字典
checkpoint = {
    'epoch': epoch + 1,                             # 当前已经训练完成的轮次数,后续恢复训练从此轮下一轮继续跑
    'model_state_dict': model.state_dict(),         # 模型当前全部权重与偏置参数
    'optimizer_state_dict': optimizer.state_dict(), # 优化器内部状态:学习率、动量、梯度累积缓存等关键信息
    'loss': train_loss,                             # 当前轮次训练集平均损失值
    'best_acc': best_acc,                           # 训练至今测试集上最优准确率,恢复后继续以此为基准保存最优模型
    'history': history,                             # 全程loss、准确率日志,恢复训练后可接续绘制训练曲线
}

# 将整套断点数据保存至本地 checkpoint.pth 文件
# 优势:中断断电、程序意外终止后,加载该文件可以无缝接着当前进度继续训练,不用从头开始训练
torch.save(checkpoint, "checkpoint.pth")
print("Checkpoint 已保存")

加载 Checkpoint 继续训练

# ---------------------- 断点加载初始化配置 ----------------------
# 初始起始训练轮次:默认从第0轮(第一轮)开始训练
start_epoch = 0
# 全局最优测试准确率初始值
best_acc = 0.0
# 初始化训练指标记录容器,存放每轮的训练/测试损失、准确率
history = {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}

# 判断本地是否存在断点保存文件,决定断点续训还是从零训练
if os.path.exists("checkpoint.pth"):
    # 读取本地保存的完整训练断点快照
    checkpoint = torch.load("checkpoint.pth")
    # 恢复网络模型的权重参数
    model.load_state_dict(checkpoint['model_state_dict'])
    # 恢复优化器内部状态(动量、自适应学习率、梯度缓存等),保证接续训练时优化策略不重置
    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
    # 赋值上次中断时已经跑完的轮次,后续循环从下一轮继续训练
    start_epoch = checkpoint['epoch']
    # 读取断点里记录的历史最优准确率;字典无该key时默认赋值0.0,防止报错
    best_acc = checkpoint.get('best_acc', 0.0)
    # 读取历史训练指标记录,无该key则沿用初始化空列表,保证绘图曲线不会断裂
    history = checkpoint.get('history', history)
    print(f"从 Epoch {start_epoch} 恢复训练")
else:
    # 不存在断点文件,没有历史训练数据,全新开始训练
    print("从头开始训练")

# ---------------------- 正式训练循环 ----------------------
# 不再固定从0开始循环,从断点记录的 start_epoch 接续执行至总轮数 num_epochs
for epoch in range(start_epoch, num_epochs):
    ...

📌 优化器也有状态?

是的!Adam 优化器内部会维护每个参数的"动量"信息。如果只保存模型参数不保存优化器状态,续训时需要重新"积累动量",刚开始几步可能会跳得比较大。

推理——让模型干活

推理单张图片

# 装饰器:关闭本轮推理过程中的梯度计算
# 推理阶段不需要反向传播更新参数,关闭梯度可节约显存、加快运算速度,杜绝多余张量占用内存
@torch.no_grad()
def predict_single_image(model, image_tensor):
    """
    对单张手写数字图片执行模型推理预测
    Args:
        model: 训练完成、加载好权重的CNN模型实例
        image_tensor: 预处理归一化完毕的图片张量,原始形状 (1, 28, 28),单通道28×28灰度图
    Returns:
        predicted_class: 模型判定的数字类别,取值 0~9
        probabilities: 长度为10的一维张量,对应0~9每个数字的预测概率
    """
    # 将模型切换至评估推理模式,禁用训练专属逻辑
    model.eval()

    # 扩充batch批量维度:原始张量 (通道, H, W) = (1, 28, 28)
    # 网络输入要求格式 [batch_size, channel, height, width]
    # unsqueeze(0) 在第0维新增批量维度,变为 (1, 1, 28, 28),适配网络输入规范
    image_tensor = image_tensor.unsqueeze(0)

    # 模型前向传播,输出原始得分logits,shape=(1,10)
    # 1代表单张样本,10对应0-9十个类别的未归一化预测分值
    outputs = model(image_tensor)
    # 使用softmax对logits做归一化处理,把原始分值转换成总和为1的概率分布
    probabilities = torch.softmax(outputs, dim=1)
    # 在类别维度上取分值最大的下标,即为预测类别;item()将张量数值转为普通Python整数
    predicted_class = outputs.argmax(dim=1).item()

    return predicted_class, probabilities[0]

推理外部图片(比如你用画图工具写的数字)

# 导入PIL图像处理库,用于本地图片读取、格式转换、尺寸缩放
from PIL import Image

def predict_from_file(model, image_path):
    """
    读取本地手写数字图片文件,完成预处理后调用模型进行识别预测
    Args:
        model: 已加载权重、处于就绪状态的MNIST训练好CNN模型
        image_path: 本地图片文件路径
    Returns:
        predicted_class: 模型识别出的手写数字 0~9
    """
    # ========== 步骤1:读取图片并适配MNIST标准输入尺寸与色彩模式 ==========
    image = Image.open(image_path).convert('L')  # 打开图片,convert('L') 转为单通道灰度图,和MNIST数据集格式统一
    image = image.resize((28, 28))              # 将图片强制缩放至28×28像素,匹配模型训练时输入尺寸

    # ========== 步骤2:执行和训练时完全一致的数据预处理 ==========
    transform = transforms.Compose([
        transforms.ToTensor(),                          # PIL图片转为张量,像素值从0~255映射到0~1
        transforms.Normalize((0.1307,), (0.3081,))     # 使用MNIST数据集全局均值、标准差标准化,和训练预处理保持一致
    ])
    image_tensor = transform(image)  # 得到符合模型输入规范的张量

    # ========== 步骤3:调用封装好的单图推理函数获取预测结果 ==========
    predicted_class, probs = predict_single_image(model, image_tensor)

    # ========== 控制台可视化输出预测信息 ==========
    print(f"预测结果:{predicted_class}")
    # 遍历0~9十个数字类别
    for i in range(10):
        # 根据当前类别的概率值生成字符进度条,概率越高色块越长
        bar = "█" * int(probs[i].item() * 50)
        # 打印类别编号、精确概率、可视化进度条
        print(f"  数字 {i}: {probs[i].item():.4f} {bar}")

    return predicted_class

# 调用示例:读取本地my_handwritten_5.png手写数字图片,执行整套识别流程
predict_from_file(model, "my_handwritten_5.png")

完整推理脚本

下面是一个独立的推理脚本,即使没有训练代码也能直接运行。你只需要有一个保存好的模型文件。

"""
独立推理脚本:加载已有模型,对图片做预测
用法:
    python infer.py --image my_digit.png          # 预测单张
    python infer.py --dir ./test_images/          # 批量预测
"""
import torch
import torch.nn as nn
from torchvision import transforms
from PIL import Image
import argparse
import os, glob


# 网络定义(必须和训练时一模一样)
class MNIST_CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
        self.pool1 = nn.MaxPool2d(2)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.pool2 = nn.MaxPool2d(2)
        self.fc1 = nn.Linear(64 * 7 * 7, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = torch.relu(self.conv1(x))
        x = self.pool1(x)
        x = torch.relu(self.conv2(x))
        x = self.pool2(x)
        x = x.view(x.size(0), -1)
        x = torch.relu(self.fc1(x))
        x = self.fc2(x)
        return x


def load_model(model_path):
    """加载训练好的模型"""
    model = MNIST_CNN()
    model.load_state_dict(torch.load(model_path, map_location='cpu'))
    model.eval()
    print(f"模型已加载:{model_path}")
    return model


def preprocess_image(image_path):
    """预处理图片:转灰度、缩放、归一化"""
    image = Image.open(image_path).convert('L')
    image = image.resize((28, 28))

    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    return transform(image).unsqueeze(0)  # 加 batch 维度


@torch.no_grad()
def predict(model, image_path):
    """预测单张图片并打印结果"""
    tensor = preprocess_image(image_path)
    outputs = model(tensor)
    probs = torch.softmax(outputs, dim=1)[0]
    pred = outputs.argmax(dim=1).item()

    print(f"\n图片: {image_path}")
    print(f"预测结果: {pred}(置信度: {probs[pred].item()*100:.1f}%)")

    # 显示 Top 3
    top3 = probs.topk(3)
    print("Top 3 可能性:")
    for i in range(3):
        digit = top3.indices[i].item()
        conf = top3.values[i].item()
        print(f"  {i+1}. 数字 {digit}: {conf*100:.1f}%")

    return pred


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument('--model', default='best_model.pth', help='模型文件路径')
    parser.add_argument('--image', help='单张图片路径')
    parser.add_argument('--dir', help='批量推理目录')
    args = parser.parse_args()

    model = load_model(args.model)

    if args.image:
        predict(model, args.image)
    elif args.dir:
        images = glob.glob(os.path.join(args.dir, '*.png')) + \
                 glob.glob(os.path.join(args.dir, '*.jpg'))
        for img in sorted(images):
            predict(model, img)
    else:
        print("请指定 --image 或 --dir 参数")

从 MNIST 测试集抽取样本当推理图片

你可能没有手写数字图片。没关系,从测试集中抽几张保存出来就行了:

def export_test_samples(num_per_class=2, output_dir='./examples'):
    """从测试集导出样本图片用于推理测试"""
    os.makedirs(output_dir, exist_ok=True)
    raw_dataset = datasets.MNIST("./data", train=False, download=True,
                                 transform=transforms.ToTensor())
    collected = {i: [] for i in range(10)}

    for img, label in raw_dataset:
        label_int = int(label)
        if len(collected[label_int]) < num_per_class:
            collected[label_int].append(img)
        if all(len(v) >= num_per_class for v in collected.values()):
            break

    for digit, imgs in collected.items():
        for j, img in enumerate(imgs):
            # 保存为 PNG
            from torchvision.utils import save_image
            save_image(img, f"{output_dir}/digit_{digit}_sample{j}.png")
            print(f"已保存: digit_{digit}_sample{j}.png")

本章小结

恭喜你!你已经走完了深度学习的一个完整闭环:

数据准备 → 构建网络 → 训练 → 验证 → 可视化 → 保存模型 → 推理预测

现在你可以:

  • ✅ 用 torch.save(model.state_dict(), ...)保存模型

  • ✅ 用 model.load_state_dict(torch.load(...))加载模型

  • ✅ 支持断点续训(保存/恢复 checkpoint)

  • ✅ 写独立推理脚本对新图片做预测

点击查看完整的 保存、加载、推理、断点续训 示例代码。

 

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号