深度学习入门:模型保存、加载与学习率调整

前言:上一篇我们学习了数据预处理与自定义数据集,把图片数据组织成了 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
加载 TorchScripttorch.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)

注意事项

要点说明
保存最佳模型不要保存最后一个,而是保存表现最好的
加载前需 evalmodel.eval() 切换到评估模式
参数名检查model.state_dict().keys() 可验证模型结构
两种保存方式训练时用 state_dict,部署时用 TorchScript
调度器参数patience 不要太小,避免学习率过早降低
调度器调用ReduceLROnPlateau 需要传入监控指标(如 loss)

系列直达

更多推荐