学会保存训练好的模型,下次不用重新训练
学会加载已保存的模型,直接用它做预测
支持断点续训——训练中断后可以从上次停下的地方继续
写一个独立推理脚本,对任意手写图片做出预测
这一章学完,你就拥有了一个真正 "能用" 的 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 = {
'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 已保存")# ---------------------- 断点加载初始化配置 ----------------------
# 初始起始训练轮次:默认从第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 参数")你可能没有手写数字图片。没关系,从测试集中抽几张保存出来就行了:
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)
✅ 写独立推理脚本对新图片做预测
点击查看完整的 保存、加载、推理、断点续训 示例代码。