深度学习入门:模型保存、加载与学习率调整
2026/9/16 7:15:47 网站建设 项目流程

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

前言:上一篇我们学习了数据预处理与自定义数据集,把图片数据组织成了 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# 记录历史最佳准确率,放在训练循环外deftest(dataloader,model,loss_fn):globalbest_acc# 声明使用全局变量size=len(dataloader.dataset)num_batches=len(dataloader)model.eval()test_loss,correct=0,0withtorch.no_grad():forX,yindataloader: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/=sizeprint(f"Test result: \n Accuracy:{(100*correct)}%, Avg loss:{test_loss}")# 如果当前模型优于历史最佳,则保存ifcorrect>best_acc:best_acc=correctprint(model.state_dict().keys())# 打印所有参数名torch.save(model.state_dict(),'food_cnn_weights.pth')# 保存权重script_model=torch.jit.script(model)# 转为 TorchScripttorch.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 调整一次学习率。

deftrain(dataloader,model,loss_fn,optimizer):model.train()batch_size_num=1forX,yindataloader: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 预测代码

importtorchimportnumpyasnpfromtorchimportnnfromtorch.utils.dataimportDataset,DataLoaderfromPILimportImagefromtorchvisionimporttransforms device="cuda"iftorch.cuda.is_available()else"mps"iftorch.backends.mps.is_available()else"cpu"# ==================== 定义模型结构(加载 state_dict 时需要)====================classCNN(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)defforward(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)returnoutput# ==================== 加载模型 ====================# 方式一:加载 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])]),}classFoodDataset(Dataset):def__init__(self,file_path,transform=None):self.imgs=[]self.labels=[]self.transform=transformwithopen(file_path)asf:samples=[x.strip().split(' ')forxinf.readlines()]forimg_path,labelinsamples:self.imgs.append(img_path)self.labels.append(label)def__len__(self):returnlen(self.imgs)def__getitem__(self,idx):image=Image.open(self.imgs[idx])ifself.transform:image=self.transform(image)label=torch.from_numpy(np.array(self.labels[idx],dtype=np.int64))returnimage,label test_data=FoodDataset(file_path='./test.txt',transform=data_transforms['valid'])test_dataloader=DataLoader(test_data,batch_size=1,shuffle=True)# ==================== 批量预测 ====================deftest_true(dataloader,model):"""返回所有样本的预测值和真实值"""result=[]labels=[]withtorch.no_grad():forX,yindataloader:X,y=X.to(device),y.to(device)pred=model.forward(X)result.append(pred.argmax(1).item())labels.append(y.item())returnresult,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)

系列直达

  • 上篇:深度学习入门:数据预处理与自定义数据集
  • 本篇:深度学习入门:模型保存、加载与学习率调整(本文)
  • 下篇:敬请期待

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询