深度学习入门:模型保存、加载与学习率调整
深度学习入门:模型保存、加载与学习率调整
前言:上一篇我们学习了数据预处理与自定义数据集,把图片数据组织成了 PyTorch 可以训练的格式。本篇我们将学习模型训练完成后的保存与加载,以及学习率动态调整策略。训练一个好的模型往往需要很长时间,把训练好的模型保存下来,下次直接加载使用,是实际项目中必不可少的环节。
目录
- 一、为什么要保存模型
- 二、两种保存方式
- 三、保存最佳模型
- 四、学习率动态调整
- 五、加载模型并预测
- 六、总结
一、为什么要保存模型
深度学习模型训练时间长,动辄几小时甚至几天。如果每次使用都要重新训练,效率极低。保存模型的好处:
| 好处 | 说明 |
|---|---|
| 省时 | 一次训练,多次使用 |
| 可复用 | 部署到服务器、嵌入式设备 |
| 可分享 | 把训练好的模型发给别人 |
| 可恢复 | 训练中断后从保存点继续 |
二、两种保存方式
PyTorch 提供两种模型保存方式:
| 方式 | 保存内容 | 特点 |
|---|---|---|
| state_dict | 只保存参数权重 | 体积小,需要模型类才能加载 |
| torch.jit.script | 保存完整模型(含结构) | 可直接加载推理,无需定义模型类 |
2.1 方式一:保存 state_dict
torch.save(model.state_dict(), 'food_cnn_weights.pth')
特点:
- 只保存权重参数,不保存模型结构
- 加载时需要先实例化 CNN 类,再加载权重
- 文件较小,适合训练阶段保存
2.2 方式二:保存完整模型(TorchScript)
script_model = torch.jit.script(model)
torch.jit.save(script_model, 'food_cnn_script.pth')
特点:
- 保存完整模型结构和参数
- 加载时不需要定义 CNN 类,可直接加载推理
- 适合部署到生产环境
三、保存最佳模型
在实际训练中,我们希望保存表现最好的那一版模型,而不是最后一版。具体做法:每次测试时比较当前准确率与历史最佳,如果更好就保存。
3.1 修改 test 函数
best_acc = 0 # 记录历史最佳准确率,放在训练循环外
def test(dataloader, model, loss_fn):
global best_acc # 声明使用全局变量
size = len(dataloader.dataset)
num_batches = len(dataloader)
model.eval()
test_loss, correct = 0, 0
with torch.no_grad():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
test_loss += loss_fn(pred, y).item()
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
test_loss /= num_batches
correct /= size
print(f"Test result: \n Accuracy: {(100 * correct)}%, Avg loss: {test_loss}")
# 如果当前模型优于历史最佳,则保存
if correct > best_acc:
best_acc = correct
print(model.state_dict().keys()) # 打印所有参数名
torch.save(model.state_dict(), 'food_cnn_weights.pth') # 保存权重
script_model = torch.jit.script(model) # 转为 TorchScript
torch.jit.save(script_model, 'food_cnn_script.pth') # 保存完整模型
3.2 保存逻辑说明
| 步骤 | 说明 |
|---|---|
| 对比准确率 | 当前准确率 > 历史最佳才保存 |
| 更新最佳值 | 保存成功后更新 best_acc |
| 打印参数名 | model.state_dict().keys() 可用于确认模型结构 |
| 保存两种格式 | 同时保存权重和完整模型,兼顾灵活性和部署 |
四、学习率动态调整
4.1 为什么需要调整学习率
学习率是深度学习最重要的超参数之一。常用的学习率有 0.1、0.01、0.001 等,学习率越大权重更新越快:
- 学习率太大:训练不稳定,损失震荡
- 学习率太小:收敛太慢,训练时间长
- 固定学习率:后期难以精细收敛
理想的做法是:训练初期用较大学习率快速收敛,训练后期用较小学习率精细调整,从而更好地收敛到最优解。
4.2 PyTorch 的三种调整方法
PyTorch 通过 torch.optim.lr_scheduler 接口实现学习率调整,提供三种方法:
| 方法 | 说明 | 代表调度器 |
|---|---|---|
| 有序调整 | 按预设的 epoch 规则调整 | StepLR、MultiStepLR、ExponentialLR、CosineAnnealingLR |
| 自适应调整 | 根据训练指标(loss、accuracy)伺机调整 | ReduceLROnPlateau |
| 自定义调整 | 通过自定义 lambda 函数调整 | LambdaLR |
4.3 有序调整
StepLR(等间隔调整)
每隔固定的 epoch 数,学习率乘以衰减系数。
scheduler = torch.optim.lr_scheduler.StepLR(
optimizer,
step_size=30, # 每 30 个 epoch 调整一次
gamma=0.1 # 学习率乘以 0.1
)
| 参数 | 说明 |
|---|---|
step_size | 学习率下降间隔数(单位:epoch) |
gamma | 学习率调整倍数,默认为 0.1 |
MultiStepLR(多间隔调整)
在指定的多个 epoch 处调整学习率。
scheduler = torch.optim.lr_scheduler.MultiStepLR(
optimizer,
milestones=[10, 30, 80], # 在第 10、30、80 个 epoch 调整
gamma=0.1
)
ExponentialLR(指数衰减)
学习率按指数规律衰减。
scheduler = torch.optim.lr_scheduler.ExponentialLR(
optimizer,
gamma=0.9 # 每个 epoch 学习率乘以 0.9
)
CosineAnnealingLR(余弦退火)
学习率按余弦函数曲线变化,先下降再上升。
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=50, # 学习率下降到最小值的 epoch 数
eta_min=0 # 学习率的最小值
)
4.4 自适应调整
ReduceLROnPlateau(根据指标调整)
当监测的指标不再改善时,自动降低学习率。这是本案例使用的调度器。
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer,
mode="min", # 监控指标是越小越好(如 loss),监控 acc 时用 "max"
factor=0.1, # 学习率衰减系数
patience=10, # 连续 10 次没有改善才降低学习率
verbose=False, # 是否打印日志
threshold=0.0001, # 改善阈值
threshold_mode='rel', # 相对变化,新值 ≤ 旧值 × (1-threshold) 才算改善
cooldown=0, # 降低学习率后冷却多少轮
min_lr=0, # 学习率下限
eps=1e-08 # 学习率最小变化量
)
| 参数 | 说明 |
|---|---|
mode | "min" 表示指标越小越好(如 loss),"max" 表示越大越好(如 acc) |
factor | 学习率衰减系数,常用 0.1 |
patience | 容忍多少次没改善后再降低学习率 |
threshold | 判定“有改善”的最小变化量 |
cooldown | 降低学习率后的冷却期 |
min_lr | 学习率的下限 |
4.5 本案例的使用方式
本案例的数据量较小,训练集只有几百张图片,batch 数量少,因此将 scheduler.step() 放在 train 的 batch 循环内,每个 batch 结束后根据当前 loss 调整一次学习率。
def train(dataloader, model, loss_fn, optimizer):
model.train()
batch_size_num = 1
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
loss = loss_fn(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_value = loss.item()
scheduler.step(loss_value) # 每个 batch 结束调用一次
print(f"loss: {loss_value:>7f} [number:{batch_size_num}]")
batch_size_num += 1
说明:
ReduceLROnPlateau通常是按 epoch 调用,但本案例数据量小、batch 数量少,放在 batch 内调用也完全可以跑通,实现简单。
4.6 各调度器对比
| 调度器 | 调整方式 | 是否需要传入指标 | 适用场景 |
|---|---|---|---|
| StepLR | 等间隔调整 | 否 | 训练轮数已知 |
| MultiStepLR | 多间隔调整 | 否 | 关键节点手动控制 |
| ExponentialLR | 指数衰减 | 否 | 平滑衰减 |
| CosineAnnealingLR | 余弦退火 | 否 | 需要周期性探索 |
| ReduceLROnPlateau | 自适应调整 | 是 | 无法预估训练轮数 |
| LambdaLR | 自定义调整 | 否 | 特殊需求 |
五、加载模型并预测
模型保存后,就可以在需要时加载使用。两种保存方式对应两种加载方式。
5.1 两种加载方式对比
| 方式 | 是否需要 CNN 类 | 适用场景 |
|---|---|---|
load_state_dict | 需要 | 训练时、修改模型结构 |
torch.jit.load | 不需要 | 部署、推理 |
5.2 加载 state_dict 模型
需要先实例化 CNN 类,再加载权重:
m1 = CNN() # 先创建模型对象
m1.load_state_dict(torch.load('food_cnn_weights.pth')) # 加载权重
m1.eval() # 切换到评估模式
5.3 加载 TorchScript 模型
不需要定义 CNN 类,直接加载:
m2 = torch.jit.load('food_cnn_script.pth') # 直接加载完整模型
m2.eval()
5.4 预测代码
import torch
import numpy as np
from torch import nn
from torch.utils.data import Dataset, DataLoader
from PIL import Image
from torchvision import transforms
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
# ==================== 定义模型结构(加载 state_dict 时需要)====================
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(in_channels=3, out_channels=16, kernel_size=5, stride=1, padding=2),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2),
)
self.conv2 = nn.Sequential(
nn.Conv2d(16, 32, 5, 1, 2),
nn.ReLU(),
nn.Conv2d(32, 32, 5, 1, 2),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.conv3 = nn.Sequential(
nn.Conv2d(32, 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)
output = self.out(x)
return output
# ==================== 加载模型 ====================
# 方式一:加载 state_dict(需要 CNN 类)
m1 = CNN()
m1.load_state_dict(torch.load('food_cnn_weights.pth'))
m1.eval()
# 方式二:加载 TorchScript 模型(不需要 CNN 类)
m2 = torch.jit.load('food_cnn_script.pth')
m2.eval()
# ==================== 准备测试数据 ====================
data_transforms = {
'valid':
transforms.Compose([
transforms.Resize((256, 256)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
]),
}
class FoodDataset(Dataset):
def __init__(self, file_path, transform=None):
self.imgs = []
self.labels = []
self.transform = transform
with open(file_path) as f:
samples = [x.strip().split(' ') for x in f.readlines()]
for img_path, label in samples:
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 = FoodDataset(file_path='./test.txt', transform=data_transforms['valid'])
test_dataloader = DataLoader(test_data, batch_size=1, shuffle=True)
# ==================== 批量预测 ====================
def test_true(dataloader, model):
"""返回所有样本的预测值和真实值"""
result = []
labels = []
with torch.no_grad():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
result.append(pred.argmax(1).item())
labels.append(y.item())
return result, labels
# 使用 m1(state_dict 加载的模型)
result1, labels1 = test_true(test_dataloader, m1)
print('预测值1:\t', result1)
print('真实值1:\t', labels1)
# 使用 m2(TorchScript 加载的模型)
result2, labels2 = test_true(test_dataloader, m2)
print('预测值2:\t', result2)
print('真实值2:\t', labels2)
5.5 输出示例
预测值1: [9, 16, 19, 16, 17, 8, 3, 8, ...]
真实值1: [6, 16, 13, 1, 17, 9, 13, 5, ...]
预测值2: [11, 11, 19, 3, 8, 3, 3, 11, ...]
真实值2: [18, 2, 18, 3, 5, 13, 12, 16, ...]
通过对比预测值和真实值,可以直观验证模型的效果。
六、总结
核心知识点速查
| 知识点 | 关键概念 |
|---|---|
| state_dict 保存 | torch.save(model.state_dict(), 'food_cnn_weights.pth') |
| TorchScript 保存 | torch.jit.save(torch.jit.script(model), 'food_cnn_script.pth') |
| 保存最佳模型 | 比较准确率,高于历史最佳才保存 |
| 学习率调度器 | ReduceLROnPlateau 自动降低学习率 |
| 加载 state_dict | 需先实例化 CNN 类,再 load_state_dict |
| 加载 TorchScript | torch.jit.load() 直接加载,无需 CNN 类 |
核心 API 一览
| 用途 | 对应方法 |
|---|---|
| 保存权重 | torch.save(model.state_dict(), path) |
| 加载权重 | model.load_state_dict(torch.load(path)) |
| 保存完整模型 | torch.jit.save(torch.jit.script(model), path) |
| 加载完整模型 | torch.jit.load(path) |
| 学习率调度 | torch.optim.lr_scheduler.ReduceLROnPlateau() |
| 调度器更新 | scheduler.step(metric) |
注意事项
| 要点 | 说明 |
|---|---|
| 保存最佳模型 | 不要保存最后一个,而是保存表现最好的 |
| 加载前需 eval | model.eval() 切换到评估模式 |
| 参数名检查 | model.state_dict().keys() 可验证模型结构 |
| 两种保存方式 | 训练时用 state_dict,部署时用 TorchScript |
| 调度器参数 | patience 不要太小,避免学习率过早降低 |
| 调度器调用 | ReduceLROnPlateau 需要传入监控指标(如 loss) |
系列直达
- 上篇:深度学习入门:数据预处理与自定义数据集
- 本篇:深度学习入门:模型保存、加载与学习率调整(本文)
- 下篇:敬请期待
更多推荐



所有评论(0)