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

相关新闻

秋叶大佬sd出现Error: tuple indices must be integers or slices, not float报错

秋叶大佬sd出现Error: tuple indices must be integers or slices, not float报错

找到webui包下的config.json改成config.json.bak修正:如果改成bak后会影响功能其实这个报错是因为clip skip是小数的原因改成整数即可

2026/9/16 7:15:57 阅读更多 →
IntelliJ IDEA社区版完全指南:免费开源IDE的下载、配置、插件与故障排查

IntelliJ IDEA社区版完全指南:免费开源IDE的下载、配置、插件与故障排查

很多人对IDEA的印象还停留在"这是个收费软件",于是网上一搜全是各种非正规激活渠道、注册码心得。这里得先说清楚一个最容易被忽略的事实:JetBrains官方一直维护着一个完全免费、开源、无功能限制的IntelliJ IDEA Community Edition&#xff0…

2026/9/16 7:14:56 阅读更多 →
Delphi 12.3集成HTML Component Library v4.6实战:安装、配置与调优

Delphi 12.3集成HTML Component Library v4.6实战:安装、配置与调优

简介:HTML Component Library v4.6 是一套针对 Delphi 平台的 HTML 渲染控件,主要解决 VCL 与 FMX 程序中嵌入网页、解析 HTML/CSS、处理表单交互等问题,支持从 Delphi 5 到 Delphi 11 Alexandria 的多个版本,适合中高级桌面应用开…

2026/9/16 7:14:56 阅读更多 →

最新新闻

基于STM32构建城市路灯节能控制系统

基于STM32构建城市路灯节能控制系统

城市路灯节能控制系统 绪论系统方案设计 系统总体架构系统功能主要模块选型 硬件设计 主控与电源传感器接口路灯驱动系统接线总表 软件设计 主程序流程调光策略太阳能追光平台运维 系统实现与测试 测试场景测试结果 总结 城市路灯节能控制系统 绪论 城市道路照明在传统定时控…

2026/9/16 8:01:15 阅读更多 →
跨国调研的样本质量怎么控?2026六家全球样本服务质控对比

跨国调研的样本质量怎么控?2026六家全球样本服务质控对比

测评说明:本文为调研从业者实测记录,记录问卷星样本服务、问卷网样本服务、51 调查、云调查、集思网等2026年六家线上样本渠道公开质控相关功能参数,面向跨国调研场景,记录各渠道可支持能力。仅陈述客观功能与规则,不作…

2026/9/16 8:01:15 阅读更多 →
Windows环境下单机版DM8安装

Windows环境下单机版DM8安装

3.1 运行镜像文件3.2 选择语言时区3.3 进入安装向导3.4 接受许可证协议3.5 Key文件配置测试版直接跳过,可免费使用一年3.6 选择安装组件选择典型安装,即包含所有组件3.7 选择安装目录安装到指定目录3.8 确认安装配置3.9 安装3.10 初始化数据库3.11 创建实…

2026/9/16 8:01:15 阅读更多 →
虚实融合数字孪生实验室:四大虚拟仿真平台重构嵌入式教学

虚实融合数字孪生实验室:四大虚拟仿真平台重构嵌入式教学

做了十年嵌入式教学和实验室建设相关工作,我见过太多高校实验室从"重金采购"到"吃灰闲置"的完整周期。很多学校花大价钱建了嵌入式实验室、物联网实验室、AI实验室,结果设备更新跟不上技术迭代,学生学的内容和产业需求严…

2026/9/16 8:01:15 阅读更多 →
船舶摇摆台厂家怎么选?先看这四类出厂检测报告和定型试验能力

船舶摇摆台厂家怎么选?先看这四类出厂检测报告和定型试验能力

摇摆台采购的“验收盲区”做过船载设备环境试验的项目采购和测试工程师都清楚一个现实:摇摆台买回来,验收环节往往是走过场。设备到货,通电试运行,看台面能动、能摇,就签收。但真正到了船载设备定型试验或者可靠性考核…

2026/9/16 8:01:15 阅读更多 →
PPO算法在无人机三维路径规划中的应用与实践

PPO算法在无人机三维路径规划中的应用与实践

1. 项目概述:无人机三维路径规划与PPO算法在无人机自主导航领域,三维路径规划一直是核心挑战之一。传统方法如A*、RRT等算法虽然成熟,但在动态复杂环境中往往表现僵硬。这个项目采用近端策略优化(PPO)这一强化学习算法…

2026/9/16 8:00:15 阅读更多 →

日新闻

嵌入式三大高薪赛道:车规功能安全、RISC-V固件架构、边缘AI部署

嵌入式三大高薪赛道:车规功能安全、RISC-V固件架构、边缘AI部署

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

2026/9/16 0:00:51 阅读更多 →
IoT-For-Beginners 智能语音计时器:Wio Terminal 基于 DMAC 与 Flash 的音频采集实战

IoT-For-Beginners 智能语音计时器:Wio Terminal 基于 DMAC 与 Flash 的音频采集实战

IoT-For-Beginners 智能语音计时器:Wio Terminal 基于 DMAC 与 Flash 的音频采集实战 【免费下载链接】IoT-For-Beginners 12 Weeks, 24 Lessons, IoT for All! 项目地址: https://gitcode.com/GitHub_Trending/io/IoT-For-Beginners 本指南聚焦 GitHub Tren…

2026/9/16 0:01:52 阅读更多 →
基于MATLAB的CRI显色指数计算:从SPD光谱到Ra的完整流程

基于MATLAB的CRI显色指数计算:从SPD光谱到Ra的完整流程

简介:针对照明设计与光学研究中的光谱功率分布(SPD)与显色性指数(CRI)计算需求,这套MATLAB程序为照明工程师、LED研发人员及光学专业学生提供了轻量工具。代码通过解析光谱测量数据,自动完成波长…

2026/9/16 0:01:52 阅读更多 →

周新闻

AI SDK Harness 依赖更新指南:掌握 harness 包 SDK 依赖的升级、桥接同步与一致性校验

AI SDK Harness 依赖更新指南:掌握 harness 包 SDK 依赖的升级、桥接同步与一致性校验

AI SDK Harness 依赖更新指南:掌握 harness 包 SDK 依赖的升级、桥接同步与一致性校验 【免费下载链接】ai The AI Toolkit for TypeScript. From the creators of Next.js, the AI SDK is a free open-source library for building AI-powered applications and ag…

2026/9/15 12:27:42 阅读更多 →
Refine v5 Ant Design NumberField 组件实战:基于 Intl 的本地化数字格式化

Refine v5 Ant Design NumberField 组件实战:基于 Intl 的本地化数字格式化

Refine v5 Ant Design NumberField 组件实战:基于 Intl 的本地化数字格式化 【免费下载链接】refine A React Framework for building internal tools, admin panels, dashboards & B2B apps with unmatched flexibility. 项目地址: https://gitcode.com/GitH…

2026/9/16 1:59:46 阅读更多 →
Flutter应用改名全指南:从Android到iOS的配置与工具实践

Flutter应用改名全指南:从Android到iOS的配置与工具实践

刚接一个外包项目时,甲方要求把工程里临时用的应用名改成正式产品名。我本来觉得“改名”这种小事,打开配置文件改一行不就完了?结果真动手才发现,Flutter项目里“应用名称”根本不是一处配置,而是一整套散落在 Androi…

2026/9/16 1:59:35 阅读更多 →

月新闻

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能分类:[AI/大模型]细分主题:AI 增强型 CI/CD 流水线自动化与 GitOps 实践:Agent 工作流、工具调用与任务拆解:从原型到生产的验收清单很多团队在尝试用大…

2026/9/15 21:40:00 阅读更多 →
容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场分类:[工程技术]细分主题:Kubernetes 生产环境运维与排障实战:可复制的项目复盘模板与决策记录大部分团队的事故复盘报告,最后都变成了躺在 Confluence 或钉…

2026/9/15 21:39:18 阅读更多 →
容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步

容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步

容器 容器化技术与镜像安全管理:核心链路应该先拆哪一步分类:[工程技术]细分主题:Docker 容器化技术与镜像安全管理:核心链路的逐步实现与关键代码取舍面对一个积累了五六年历史包袱的单体架构应用(包含 Web 接口、后台…

2026/9/15 21:40:17 阅读更多 →