深度学习入门:模型保存、加载与学习率调整
前言:上一篇我们学习了数据预处理与自定义数据集,把图片数据组织成了 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 |
| 加载 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) |
系列直达
- 上篇:深度学习入门:数据预处理与自定义数据集
- 本篇:深度学习入门:模型保存、加载与学习率调整(本文)
- 下篇:敬请期待