本节你会学到
- 说清楚「模型保存、加载和批量推理实战」解决的核心问题
- 知道它在「深度学习项目」中的位置
- 把 DataLoader、Linear 层、Tensor 这些关键词联系起来
训练好的模型不能只活在内存里。这篇讲清怎么把模型存成文件、怎么加载回来预测、怎么做批量推理,以及断点续训的完整做法。
学习路线:深度学习项目 · 第 63 课
通过回归、二维分类、表格分类和图像分类项目,把 Dataset、模型、损失、优化器、评估和推理串起来。
学完本阶段你能做到:能独立完成一个完整深度学习项目,从数据到推理全流程,并整理出可复用的项目目录结构。
推荐读法:建议顺序阅读并同步跑代码;每个项目都要形成可复用目录。
查看完整阶段 · 10 篇先看目标,再带着问题读正文。读完后用练习确认自己真的理解了。
先检查模型参数、输入 Tensor、标签 Tensor 是否在同一个设备上。常见修正是把它们统一 `.to(device)`。
因为训练是否正确经常取决于 shape、dtype、device 和梯度流,API 名字只能告诉你工具,不会保证数据契约正确。
`输入数据 → forward → loss → backward → optimizer.step()`,并知道本文主题位于这条链路的哪一环。
你花 3 小时训练了一个模型,关掉电脑就没了。下次想用还得重训。这就是为什么必须学会模型保存——训练成果要能存成文件、随时加载、批量预测。
第 50 篇讲了 nn.Module 的 state_dict() 返回模型参数。第 45 篇讲了优化器也有内部状态。今天把它们存到文件里。第 56 篇讲了 GPU 训练——保存加载时还要处理设备问题。
用
torch.save存模型参数,用load_state_dict加载。存的是state_dict(参数字典),不是整个模型对象。
torch.save(model.state_dict(), 'model.pth')
.state_dict() 返回一个字典,包含所有参数的名称和值:
print(model.state_dict().keys())
# odict_keys(['net.0.weight', 'net.0.bias', 'net.2.weight', 'net.2.bias', ...])
# 1. 先定义相同结构的模型
model = TabularClassifier(in_features=20, num_classes=4)
# 2. 加载参数
model.load_state_dict(torch.load('model.pth'))
# 3. 切到评估模式
model.eval()
关键:加载前必须先创建相同结构的模型对象——load_state_dict 只填参数,不创建模型。
with torch.no_grad():
x_new = torch.tensor([[...]]).float()
logits = model(x_new)
pred = logits.argmax(dim=1)
print(f"预测类别: {pred.item()}")
训练中断了想继续?需要存更多东西:模型参数 + 优化器状态 + 当前 epoch。
checkpoint = {
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'loss': loss.item(),
}
torch.save(checkpoint, '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'] + 1
# 继续训练
for epoch in range(start_epoch, total_epochs):
...
为什么要存优化器状态? Adam 内部维护动量($m_t$、$v_t$),如果不恢复,等于从头开始算动量,前几步更新方向会偏。
实际应用中,要对大量数据做预测。
import torch
from torch.utils.data import DataLoader
def batch_inference(model, dataset, device='cpu', batch_size=64):
"""批量推理,返回所有预测结果"""
model.eval()
model = model.to(device)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=False)
all_preds = []
all_probs = []
with torch.no_grad():
for batch_x in loader:
if isinstance(batch_x, (list, tuple)):
batch_x = batch_x[0]
batch_x = batch_x.to(device)
logits = model(batch_x)
probs = torch.softmax(logits, dim=1)
preds = logits.argmax(dim=1)
all_preds.extend(preds.cpu().numpy())
all_probs.extend(probs.cpu().numpy())
return all_preds, all_probs
# 使用
preds, probs = batch_inference(model, test_dataset, device='cuda')
print(f"前 10 个预测: {preds[:10]}")
print(f"前 10 个概率: {probs[:10]}")
关键点:
shuffle=False 保持顺序,方便对照torch.no_grad() 不建计算图,省显存model.eval() 关闭 Dropout 和 BatchNorm# GPU 上保存的模型,在 CPU 机器上加载
model.load_state_dict(torch.load('model.pth', map_location='cpu'))
map_location='cpu' 把 GPU 上的参数映射到 CPU。
model.load_state_dict(torch.load('model.pth'))
model = model.to('cuda')
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.load_state_dict(torch.load('model.pth', map_location=device))
model = model.to(device)
训练时保存验证集上表现最好的模型:
best_val_acc = 0
for epoch in range(epochs):
# 训练...
val_acc = evaluate(model, val_loader)
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), 'best_model.pth')
print(f"保存最佳模型,Val Acc={val_acc:.4f}")
# 训练完后加载最佳模型
model.load_state_dict(torch.load('best_model.pth'))
import torch
import torch.nn as nn
# === 训练阶段 ===
model = MyModel()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
criterion = nn.CrossEntropyLoss()
best_acc = 0
for epoch in range(100):
train(...)
val_acc = evaluate(...)
if val_acc > best_acc:
best_acc = val_acc
torch.save({
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'epoch': epoch,
'val_acc': val_acc,
}, 'best_checkpoint.pth')
# === 推理阶段(另一个脚本)===
model = MyModel() # 相同结构
checkpoint = torch.load('best_checkpoint.pth')
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
with torch.no_grad():
logits = model(new_data)
pred = logits.argmax(dim=1)
# ❌ 存整个模型(pickle)
torch.save(model, 'model.pth')
# ✅ 只存 state_dict
torch.save(model.state_dict(), 'model.pth')
存整个模型会绑定代码路径和类定义,换环境容易报错。存 state_dict 更通用。
# 训练时
class Model(nn.Module):
def __init__(self):
self.fc = nn.Linear(20, 10)
# 加载时改了结构
class Model(nn.Module):
def __init__(self):
self.fc = nn.Linear(30, 10) # ❌ 维度变了
model.load_state_dict(torch.load('model.pth'))
# RuntimeError: Error(s) in loading state_dict
修复:加载时的模型结构必须和保存时完全一致。
# ❌ 推理时 Dropout 还在随机
predictions = model(test_data)
# ✅
model.eval()
with torch.no_grad():
predictions = model(test_data)
练习 1:断点续训时为什么要恢复优化器的 state_dict?不恢复会怎样?
练习 2:写一个函数,接收模型路径和数据,完成单条数据的预测。
练习 3:GPU 训练保存的模型,在只有 CPU 的服务器上加载,代码怎么写?
练习 1:Adam 内部维护一阶动量 $m_t$ 和二阶动量 $v_t$,这些是基于历史梯度累积的。不恢复优化器状态,等于从头算动量,前几步的更新方向会不稳定,可能影响收敛。SGD+Momentum 同理(动量 $v_t$ 丢失)。对纯 SGD(无 momentum)影响较小。
练习 2:
def predict_single(model_path, x, model_class, in_features, num_classes):
model = model_class(in_features, num_classes)
model.load_state_dict(torch.load(model_path, map_location='cpu'))
model.eval()
with torch.no_grad():
x = torch.tensor(x).float().unsqueeze(0) # 加 batch 维
logits = model(x)
pred = logits.argmax(dim=1).item()
return pred
练习 3:
model = MyModel()
model.load_state_dict(torch.load('model.pth', map_location='cpu'))
model.eval()
# 在 CPU 上推理,不需要 .to('cuda')
state_dict,不存整个模型对象map_location下一篇进入图像世界——CNN 的卷积和池化到底在提取什么。
留言讨论