尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

深度学习入门:模型保存、加载与学习率调整 深度学习入门模型保存、加载与学习率调整前言上一篇我们学习了数据预处理与自定义数据集把图片数据组织成了 PyTorch 可以训练的格式。本篇我们将学习模型训练完成后的保存与加载以及学习率动态调整策略。训练一个好的模型往往需要很长时间把训练好的模型保存下来下次直接加载使用是实际项目中必不可少的环节。目录一、为什么要保存模型二、两种保存方式三、保存最佳模型四、学习率动态调整五、加载模型并预测六、总结一、为什么要保存模型深度学习模型训练时间长动辄几小时甚至几天。如果每次使用都要重新训练效率极低。保存模型的好处好处说明省时一次训练多次使用可复用部署到服务器、嵌入式设备可分享把训练好的模型发给别人可恢复训练中断后从保存点继续二、两种保存方式PyTorch 提供两种模型保存方式方式保存内容特点state_dict只保存参数权重体积小需要模型类才能加载torch.jit.script保存完整模型含结构可直接加载推理无需定义模型类2.1 方式一保存 state_dicttorch.save(model.state_dict(),food_cnn_weights.pth)特点只保存权重参数不保存模型结构加载时需要先实例化 CNN 类再加载权重文件较小适合训练阶段保存2.2 方式二保存完整模型TorchScriptscript_modeltorch.jit.script(model)torch.jit.save(script_model,food_cnn_script.pth)特点保存完整模型结构和参数加载时不需要定义 CNN 类可直接加载推理适合部署到生产环境三、保存最佳模型在实际训练中我们希望保存表现最好的那一版模型而不是最后一版。具体做法每次测试时比较当前准确率与历史最佳如果更好就保存。3.1 修改 test 函数best_acc0# 记录历史最佳准确率放在训练循环外deftest(dataloader,model,loss_fn):globalbest_acc# 声明使用全局变量sizelen(dataloader.dataset)num_batcheslen(dataloader)model.eval()test_loss,correct0,0withtorch.no_grad():forX,yindataloader:X,yX.to(device),y.to(device)predmodel.forward(X)test_lossloss_fn(pred,y).item()correct(pred.argmax(1)y).type(torch.float).sum().item()test_loss/num_batches correct/sizeprint(fTest result: \n Accuracy:{(100*correct)}%, Avg loss:{test_loss})# 如果当前模型优于历史最佳则保存ifcorrectbest_acc:best_acccorrectprint(model.state_dict().keys())# 打印所有参数名torch.save(model.state_dict(),food_cnn_weights.pth)# 保存权重script_modeltorch.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 函数调整LambdaLR4.3 有序调整StepLR等间隔调整每隔固定的 epoch 数学习率乘以衰减系数。schedulertorch.optim.lr_scheduler.StepLR(optimizer,step_size30,# 每 30 个 epoch 调整一次gamma0.1# 学习率乘以 0.1)参数说明step_size学习率下降间隔数单位epochgamma学习率调整倍数默认为 0.1MultiStepLR多间隔调整在指定的多个 epoch 处调整学习率。schedulertorch.optim.lr_scheduler.MultiStepLR(optimizer,milestones[10,30,80],# 在第 10、30、80 个 epoch 调整gamma0.1)ExponentialLR指数衰减学习率按指数规律衰减。schedulertorch.optim.lr_scheduler.ExponentialLR(optimizer,gamma0.9# 每个 epoch 学习率乘以 0.9)CosineAnnealingLR余弦退火学习率按余弦函数曲线变化先下降再上升。schedulertorch.optim.lr_scheduler.CosineAnnealingLR(optimizer,T_max50,# 学习率下降到最小值的 epoch 数eta_min0# 学习率的最小值)4.4 自适应调整ReduceLROnPlateau根据指标调整当监测的指标不再改善时自动降低学习率。这是本案例使用的调度器。schedulertorch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,modemin,# 监控指标是越小越好如 loss监控 acc 时用 maxfactor0.1,# 学习率衰减系数patience10,# 连续 10 次没有改善才降低学习率verboseFalse,# 是否打印日志threshold0.0001,# 改善阈值threshold_moderel,# 相对变化新值 ≤ 旧值 × (1-threshold) 才算改善cooldown0,# 降低学习率后冷却多少轮min_lr0,# 学习率下限eps1e-08# 学习率最小变化量)参数说明modemin表示指标越小越好如 lossmax表示越大越好如 accfactor学习率衰减系数常用 0.1patience容忍多少次没改善后再降低学习率threshold判定“有改善”的最小变化量cooldown降低学习率后的冷却期min_lr学习率的下限4.5 本案例的使用方式本案例的数据量较小训练集只有几百张图片batch 数量少因此将scheduler.step()放在train的 batch 循环内每个 batch 结束后根据当前 loss 调整一次学习率。deftrain(dataloader,model,loss_fn,optimizer):model.train()batch_size_num1forX,yindataloader:X,yX.to(device),y.to(device)predmodel.forward(X)lossloss_fn(pred,y)optimizer.zero_grad()loss.backward()optimizer.step()loss_valueloss.item()scheduler.step(loss_value)# 每个 batch 结束调用一次print(floss:{loss_value:7f}[number:{batch_size_num}])batch_size_num1说明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 类再加载权重m1CNN()# 先创建模型对象m1.load_state_dict(torch.load(food_cnn_weights.pth))# 加载权重m1.eval()# 切换到评估模式5.3 加载 TorchScript 模型不需要定义 CNN 类直接加载m2torch.jit.load(food_cnn_script.pth)# 直接加载完整模型m2.eval()5.4 预测代码importtorchimportnumpyasnpfromtorchimportnnfromtorch.utils.dataimportDataset,DataLoaderfromPILimportImagefromtorchvisionimporttransforms devicecudaiftorch.cuda.is_available()elsempsiftorch.backends.mps.is_available()elsecpu# 定义模型结构加载 state_dict 时需要classCNN(nn.Module):def__init__(self):super(CNN,self).__init__()self.conv1nn.Sequential(nn.Conv2d(in_channels3,out_channels16,kernel_size5,stride1,padding2),nn.ReLU(),nn.MaxPool2d(kernel_size2),)self.conv2nn.Sequential(nn.Conv2d(16,32,5,1,2),nn.ReLU(),nn.Conv2d(32,32,5,1,2),nn.ReLU(),nn.MaxPool2d(2),)self.conv3nn.Sequential(nn.Conv2d(32,128,5,1,2),nn.ReLU(),)self.outnn.Linear(128*64*64,20)defforward(self,x):xself.conv1(x)xself.conv2(x)xself.conv3(x)xx.view(x.size(0),-1)outputself.out(x)returnoutput# 加载模型 # 方式一加载 state_dict需要 CNN 类m1CNN()m1.load_state_dict(torch.load(food_cnn_weights.pth))m1.eval()# 方式二加载 TorchScript 模型不需要 CNN 类m2torch.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,transformNone):self.imgs[]self.labels[]self.transformtransformwithopen(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):imageImage.open(self.imgs[idx])ifself.transform:imageself.transform(image)labeltorch.from_numpy(np.array(self.labels[idx],dtypenp.int64))returnimage,label test_dataFoodDataset(file_path./test.txt,transformdata_transforms[valid])test_dataloaderDataLoader(test_data,batch_size1,shuffleTrue)# 批量预测 deftest_true(dataloader,model):返回所有样本的预测值和真实值result[]labels[]withtorch.no_grad():forX,yindataloader:X,yX.to(device),y.to(device)predmodel.forward(X)result.append(pred.argmax(1).item())labels.append(y.item())returnresult,labels# 使用 m1state_dict 加载的模型result1,labels1test_true(test_dataloader,m1)print(预测值1:\t,result1)print(真实值1:\t,labels1)# 使用 m2TorchScript 加载的模型result2,labels2test_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系列直达上篇深度学习入门数据预处理与自定义数据集本篇深度学习入门模型保存、加载与学习率调整本文下篇敬请期待
返回列表