资讯动态

模型训练完成后如何保存与引用?CNN模型落地工程实践

发布时间:2026/9/9 13:13:31 来源:尧图企业网站定制
模型训练完了然后呢聊聊保存与引用的那些“隐形工程”我一直觉得很多做深度学习的朋友都有一种错觉模型训练完看到loss曲线下降、验证集精度达标就以为大功告成了。结果真到用的时候要么在换了一台机器之后怎么也加载不了要么训练好的模型发过去对方连predict都跑不起来。如果你也被这种问题折磨过那这次的内容应该能帮你省下不少折腾的时间。这篇博文想聊的就是“卷积神经网络训练模型训练完之后怎么把成果真正落地”——具体来说就是模型的保存与引用。这里的“引用”不只是代码里import一个函数那么简单它牵扯到模型文件格式、网络结构定义、依赖环境、预处理配置等一堆细节。适合谁看适合那些已经能跑通CNN训练流程但还没正经处理过“模型交付”这个环节的同学也适合准备把模型从实验环境搬到其他平台、设备上用的工程师。我会把模型保存的几种方案、背后的设计逻辑、跨环境引用时的坑以及我自己实际排查过的问题挨个说一遍。1. 保存模型之前先把这三件事想明白1.1 你保存的是“模型”还是“模型一切上下文”先说一个我反复遇到的误区。很多新手第一次保存模型就是取个名字model.pth用torch.save(model.state_dict(), model.pth)一行代码搞定。等到要用的时候发现需要重建网络结构但结构定义写在另一个脚本里于是把文件拷来拷去最后干脆在Jupyter里手动拼一个网络类结果维度都对不上直接报错。关键点在于“模型”这个词在不同场景下指的是不同的东西网络结构代码——就是那个定义了nn.Module的类比如class MyCNN(nn.Module): ...。它是一段代码不是数据文件。权重参数——也就是state_dict里那一堆Parameter和buffer张量它们才是训练出来的核心成果。预处理配置——均值、标准差、图像尺寸、归一化方式等这些不在模型文件里但推理时少了它们跑出来的结果就是错的。后处理逻辑——比如检测模型里的NMS阈值、分类模型里的类别映射表同样不保存在模型文件里但对“引用模型”的人来说必不可少。如果只保存了其中的第2项那这个“模型”是不完整的。我见过有人把state_dict单独扔给对方对方加载时因为不知道训练时用的图像尺寸是224还是256推理结果完全乱套。所以你在保存之前先要划分清楚你要交付的是一整套“可复现上下文”还是一个“可加载的权重文件”这决定了后续所有操作方式。对我个人来说只要是给别人用的模型我绝不会只交一个权重文件。最次也要附带一个README或者config.yaml里面写清输入输出、预处理参数、依赖库版本。1.2 你是要“续跑训练”还是“部署推理”这个看起来是常识但实际中总有混用的情况。续跑训练场景下你希望模型停在某个epoch的状态包括优化器的动量信息、学习率调度器的进度、随机数生成器的状态等。部署推理场景则恰恰相反这些东西一个都不需要你只需要一个前向计算能跑通的文件。用一个具体例子说。PyTorch里最朴素的保存方式是这样torch.save(model.state_dict(), cnn_weights.pth)这种方案保存的东西很轻量续跑的时候你自己重建网络结构、重建优化器然后把权重灌进去虽然可行但如果你训练中断时有一堆精心调节过的优化器状态就全丢了。而深度学习炼丹老手通常会这样保存checkpoint { epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), scheduler_state_dict: scheduler.state_dict(), best_val_acc: best_val_acc, config: config, # 超参数与预处理配置 class_names: class_names } torch.save(checkpoint, fcheckpoint_epoch{epoch}_val{best_val_acc:.4f}.pth)所以我的建议是如果你的模型可能还要继续训练那就用完整checkpoint风格如果模型已经训练完毕进入“交付/部署/给别人引用”阶段再针对推理场景做精简。不要把续跑checkpoint和部署用文件混在一起不然文件会越来越大引用时还要处理一堆无关字段。1.3 框架版本与依赖环境的一致性决定了“能不能被引用”这可能是最隐蔽、也最让人头疼的一点。同一个.pth文件在PyTorch 1.8下保存在PyTorch 2.1下加载表面上可能不报错但底层算子的默认行为可能已经有变化导致输出结果有细微差异。更怕的是你用了某个高版本才有的API保存低版本直接就不认识文件头了。至于TensorFlow那边的SavedModel格式同样有兼容性问题。h5格式相对稳定但如果你在保存时用了自定义层且自定义层没有在加载环境里注册加载时会直接报Unknown layer的错误。我的个人习惯是在训练工程根目录放一个requirements.txt并且在保存模型的同一个目录里放一个environment.yaml或者pip freeze requirements.txt的结果。这样哪怕隔了半年再回来引用模型也能清楚知道自己当时是在什么环境下训练和保存的。这不算什么高深技巧但你真遇到“模型加载不了又查不出原因”的情况时这个文件就是救命稻草。2. PyTorch体系下模型保存的“组合拳”怎么打2.1 三种常见保存方式适用的场景完全不同PyTorch生态里保存模型至少有三种主流方式。我直接用表格做一个对比先说结论再展开保存方式典型代码文件内容适用场景坑点仅state_dicttorch.save(model.state_dict(), w.pth)权重张量模型结构确定、引用方能照抄网络定义缺结构代码就废了完整模型torch.save(model, m.pth)结构权重快速保存、同环境加载结构被序列化后跨环境容易出问题完整checkpointtorch.save({model: model.state_dict(), optimizer: ..., ...}, c.pth)权重优化器配置等训练续跑、实验留档文件大引用时需自己拆字段今天很多从YOLOv5、EasyOCR、PaddleOCR这类开源项目入手的朋友接触到的大多是第三种或更复杂的方案。因为这类项目本身训练和推理分离训练脚本里保存的是一个打包好的checkpoint目录里面不止有权重还有anchor配置、类别名、训练参数等。这就意味着你在“引用”这些模型时不能只盯着模型文件本身还必须把配套的配置文件一并加载。2.2 自定义网络类与加载时的“结构代码注册”问题讲一个我实际遇到过的问题。有个项目里我定义了一个带注意力模块的CNN类放在了项目路径models/attention_cnn.py里训练时保存了state_dict。后来我把这个类重构了一下改到了networks/cnn.py类名也从AttentionCNN改成了AttnNet。结果加载旧权重时直接报错RuntimeError: Error(s) in loading state_dict for AttnNet: Missing key(s) in state_dict: features.0.weight, ... Unexpected key(s) in state_dict: features_old.0.weight, ...原因很简单state_dict保存的是每个参数张量的“键名”。键名来源于你定义网络时给每个模块起的名字。只要你的模块命名变了、层的顺序变了、卷积核数量变了键名就和旧文件对不上加载自然失败。解决这个问题的方法有三条路线保持类定义和路径不变。如果模型是给自己长期使用的尽量少改结构代码或者用__init__.py做一层稳定的接口。比如原来在models.attention_cnn里的类即使内部改逻辑也保持模块路径和类名不动。加载时做key映射兼容。PyTorch允许你在load_state_dict时传入strictFalse然后手动修改键名。但这是临时的不建议长期依赖。换一种序列化方式让结构本身不再依赖Python类。这就是ONNX、TensorRT这类模型的优势——它们把网络结构和权重一起打包成中间表示引用方不再需要你原来的代码类。CNN模型部署时走这条路线很常见。我自己的习惯是实验阶段的模型用PyTorch原生方式保存但只要是准备对外“引用”的模型我一定会导出ONNX或TensorRT格式彻底摆脱Python类版本带来的藕断丝连。2.3 checkpoint命名与版本管理的实操建议说了这么多还没聊到特别具体但很实用的“命名规范”问题。很多项目的模型目录过一段时间就会变成这样model_final.pth model_final_2.pth model_final_really.pth model_final_really_v2_use_this.pth说实话这种命名方式早晚害死你。因为模型文件不像代码你不能一眼从文件内容里看出它训练到第几个epoch、精度多少、用的是哪套配置。等你想回滚到某个指标最好的版本完全只能靠猜。我推荐的做法是保存时把关键信息直接写进文件名让文件本身变成自解释的cnn_epoch150_val91.2_lr1e-4.pth cnn_best_val91.8_epoch142.pth同时配合checkpoint里的best_val_acc字段在训练代码里做“历史最优模型”自动保存。这样不管过多久你只看文件名就能快速定位到想要的那一版。还有一个实操小技巧每次计划性地保存完模型后顺手把当前的git commit id写进checkpoint的元信息字段里。这样你引用的模型对应哪一份训练代码一条命令就能查出来。这个习惯在团队协作时尤其有用。3. 跨框架、跨平台部署时模型的“导出与引用”远不止存文件3.1 从PyTorch到ONNX把结构“固化”下来不走样很多场景下训练用的框架是PyTorch但真正部署用的环境可能是TensorRT、OpenVINO、ONNX Runtime甚至是在手机端。这时候如果你还在传.pth文件对方根本没法加载。正确做法是先把模型导出成ONNX格式变成一种“中性”的模型描述格式。来看一段最简的PyTorch转ONNX代码import torch model MyCNN() checkpoint torch.load(cnn_best.pth, map_locationcpu) model.load_state_dict(checkpoint[model_state_dict]) model.eval() dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, cnn.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version11, )这里有几个关键细节值得展开讲。第一dummy_input的形状必须和训练时的输入完全一致否则导出的模型结构可能没问题但实际推理时输入的尺寸一换就报错。第二dynamic_axes的设计很重要如果你希望在推理时支持可变batch甚至可变尺寸典型的就是YOLO类检测模型要支持不同分辨率输入就必须在这里声明动态轴如果模型固定尺寸那就无所谓。第三opset_version的选择要结合部署端的运行时版本ONNX Runtime版本较老时你选了高版本opset导出的模型可能直接加载失败。导出完成后我强烈建议用onnxruntime做一次“喂同一张输入、对比输出”的验证。这一步能暴露很多问题比如某些算子导出时没有正确映射、动态轴设错导致输出维度异常、数值精度有偏差等。不要等到部署环境里才发现模型是“废的”。3.2 模型量化与精简从GPU训练到CPU、边缘设备的引用差距在GPU上训练的CNN模型权重的数据格式通常是float32模型文件往往几百MB。真到边缘设备或者普通CPU服务器上引用时这种文件又大又慢甚至内存都撑不住。这时候就需要模型量化和精简。量化最常见的做法是把权重从float32降到int8。以ONNX Runtime的量化为例静态量化流程大致是准备一批有代表性的校准数据几百到几千张样本即可。运行模型采集每层激活值的范围min/max。根据范围把权重和激活量化到int8。导出量化后的模型再对比量化前后的推理精度。量化后模型体积能降到原来的四分之一左右推理速度在CPU上通常有明显提升但精度会有一定的损失。分类任务可能只掉0.5到1个点检测任务掉得会稍多。如果你是做边缘端部署要提前想清楚“精度能接受掉多少”这个底线。另外很多人忽略的一点是量化后的模型在引用时对输入数据的格式更敏感。原始模型输入一个float32的Tensor没问题量化模型可能期望的输入范围是0到255而不是0到1这取决于你选用的量化方案。所以保存模型时一定要把“输入预处理规则”一起记录下来。否则别人拿着你的量化模型按老办法归一化结果输出全部异常排查半天发现是预处理的问题。3.3 在另一台机器上“引用”模型到底要拷贝哪些文件这个坑我踩过不止一次。训练在GPU服务器上部署在客户的内网机器上两边完全隔离。你辛辛苦苦导出的模型拷过去了却总是跑不起来。回顾这个过程其实很多问题出在“只带走了模型文件没带走必要的运行环境”。一次性完整的模型交付至少应该包含模型文件本身.pth或.onnx或.engine等。推理脚本或调用示例哪怕只有十几行也要写清楚怎么加载、怎么输入、怎么输出。依赖清单如requirements.txt或environment.yaml。输入预处理说明图像尺寸、通道顺序、归一化参数、是否BGR等。后处理与类别标签如class_names.txt、NMS参数等。版本和验证结果例如“在xx数据集上mAP为xx”这样的记录方便接收方核对。很多人嫌麻烦只发一个权重文件。但真遇到“模型在你机器上就是正常的到我机器上就崩了”这种问题返工成本远高于一开始多写几行文档的成本。我现在的习惯是目录结构做成这样deploy/ ├── model.onnx ├── config.yaml ├── class_names.txt ├── infer.py ├── requirements.txt └── test_input.jpg然后把整个目录打包交付。这样对方在pip install -r requirements.txt之后直接运行python infer.py就能复现整个推理过程基本不会再出什么幺蛾子。4. 引用模型时的报错排查完整链路我踩过的坑4.1 案例一RuntimeError: Error(s) in loading state_dict到底在说什么这是“引用PyTorch模型”时最常见的报错没有之一。很多朋友一看到它就慌其实它表达的意思非常直接你当前代码构造出来的网络和权重文件里保存的键名对不上。排查顺序我一般是这样第一步把报错里的信息完整看一遍。它会列出Missing key(s)和Unexpected key(s)。前者表示你当前网络里有某些参数权重文件里没有对应键名后者表示权重文件里有某些键你当前网络里找不到。两种情况的修复方向截然不同。第二步检查网络类的定义是否和训练时一致。最常见的不一致点包括卷积层的in_channels、out_channels改过了层与层之间的顺序做过调整某个模块的类名改了return的逻辑变了但结构没变。第三步如果是小白阶段还在用“别人训练好的backbone”或“公开模型权重”那么更可能出现的问题是你把ResNet18的权重加载到ResNet50上或者把ImageNet预训练分类头加载到自己的多分类头上。这种时候可以选择load_state_dict(..., strictFalse)让两边的公共部分先加载再单独处理分类头。这里也解释一下为什么有时候strictFalse也能跑通但效果很差。因为你新网络里随机初始化的层并不会因为加载了部分权重就自动变好。它只是“不报错而已”模型输出基本属于乱猜状态。所以strictFalse只能作为权宜之计不要当成万能钥匙。4.2 案例二ModuleNotFoundError: No module named models这个报错常见于你引用别人开源项目训练的模型他们保存的checkpoint是完整模型或依赖自定义模块的state_dict然后你把它放到自己的项目里加载结果找不到对应的网络定义。我当时排查这个问题的过程印象很深。把一个YOLOv5类项目的权重文件拖进自己的工程加载时报错找不到models.yolo这个模块。第一反应是装依赖然后发现YOLOv5的代码不是pip包而是一个仓库。解决的思路其实很简单要么把整个仓库作为依赖引入要么加载完权重后立刻把结构代码抽成独立文件。但这里有一个更值得注意的点你在“引用”一个长期更新的开源项目时同一个权重文件可能依赖特定版本的仓库代码。比如YOLOv5在某个commit之后改了模型定义你直接拉最新代码去加载老权重同样会报结构不匹配。所以引用此类模型时最好用项目官方release里一起发布的“模型代码”配套版本而不要混合使用。4.3 案例三模型文件本身损坏与加载异常模型文件损坏的情况比大家想象中多。常见触发因素包括网络传输中断、磁盘空间不足导致写入不完整、网盘同步冲突、git大文件被截断等。表现往往是加载时直接报EOFError或者unpickling error。排查这类问题有一个快速判定方法看文件大小是否正常。比如你训练完一个模型记录下正常的文件大小是342MB结果拷到另一台机器上变成300MB那基本可以断定文件有问题不用继续浪费时间调试代码。另一个与文件完整性相关的问题是torch.save保存时如果进程被强杀pickle过程可能只写入了一半导致文件结构不完整。所以保存模型时尽量在关键节点落盘不要用“训练到最后一次性保存”的方式。我对续跑场景的推荐是每个epoch或每N个epoch保存一次临时checkpoint训练结束后再额外导出一个最终的推理模型文件。这样即使中途崩溃也能从上个完整checkpoint继续而不是从零开始。5. 引用模型时经常被忽略的若干个实操细节5.1 训练中断时自动保存比“训练完再存”更可靠我在前面已经提到了这一点但这里想展开讲讲“如何设计自动保存”。很多训练脚本里是这样写的for epoch in range(epochs): train_one_epoch() validate() torch.save(model.state_dict(), final.pth)这段代码的致命弱点是如果训练到第80个epoch时服务器掉了前79个epoch全部白费。更合理的做法是每轮验证结束后达到优化目标就保存if val_acc best_val_acc: best_val_acc val_acc torch.save({ epoch: epoch 1, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_val_acc: best_val_acc, config: train_config, }, fbest_model_epoch{epoch1:03d}_val{best_val_acc:.4f}.pth)这样即便训练中断你手里也有一个精度最佳的模型可以直接引用或继续训练。这个习惯对动辄训练几小时、十几小时的CNN任务来说价值非常大。5.2 每个模型文件都应对应“验证脚本”而不只是“加载脚本”所谓引用不只是把权重load进来更重要的是用起来。我见过不少同事训练完保存了模型加载一段代码后只打印了一下输出的shape就以为大功告成。结果到部署阶段发现输出的类别索引和标签对不上、归一化方式用错、输入通道改成BGR后完全不可用。所以我坚持的实操习惯是保存模型文件的同一天就写一个简单的验证脚本输入一张固定图片跑一次前向推理打印预测结果并把输出和期望值记录下来。比如python infer.py --mode validate --model checkpoints/best.pth --image test_input.jpg # Expected output: class_id7, class_namecat, confidence0.9821把“期望输出”直接写在脚本注释或说明文档里。这样即使两周后再引用这个模型只要运行验证脚本马上就能确认模型文件和环境是否工作正常。这个方法省了我大量重复调试的时间。5.3 可以顺便留意的从EasyOCR、PaddleOCR这些开源工具“引用”模型时的共性现在很多实际项目里大家更常接触的不是自己从零训练的CNN而是基于开源OCR工具微调或直接复用的模型。此时“保存与引用”其实已经由框架封装好了但仍有几个共性经验模型路径和配置路径要保持一致。比如PaddleOCR的模型目录下通常会有inference.pdmodel、inference.pdiparams和inference.yml几个文件引用时缺一个都跑不起来。版本匹配很关键。不同版本的OCR框架对模型文件的组织方式可能做了调整旧模型用新框架加载有时能加载但推理结果明显不对所以最好直接使用训练时对应的框架版本。微调后的模型不要用原始预训练模型的配置文件去引用。因为微调时可能修改了类别数、字典表、输入尺寸这些信息都记录在微调后的配置中混用配置会导致输出边界或解码错乱。这些经验归结起来本质上就是在强调同一个道理模型从来不是一个孤立的文件而是一整套“代码权重配置依赖”的组合体。你保存和引用的时候只有把这一整套都考虑进去才算真正完成了CNN模型的闭环。5.4 安全与备份把“唯一一份”变成“永远有备份”最后说一个可能被很多人忽视的点。不要让自己辛苦训练出来的模型只有一份文件。我见过一个非常惨痛的案例某同事把训练好的模型只存在服务器本地结果服务器磁盘中文件被误删整个月的训练成果直接没了。虽然后来通过一些数据恢复工具找回了一部分但终究不是完整版本。我的建议是按“3-2-1”备份原则来管理模型文件保留3份副本使用2种不同介质至少有1份异地备份。放到具体场景里可以是训练服务器一份、本地移动硬盘一份、对象存储或网盘一份。模型文件往往动辄几百MB甚至一个checkpoint就到几个GB你可能觉得备份起来麻烦但相对于几天甚至几周的训练成本这点存储开销完全值得。如果工程上走的是Git也要注意不要把模型文件直接纳入普通的Git仓库。模型文件变更频繁且体积大会让仓库迅速膨胀。更合适的做法是使用专门的模型版本管理工具或者简化为“在代码里记录模型文件对应的下载链接和版本号”。这样既不会拖慢日常开发也方便回滚到某一个具体版本。对我来说模型保存与引用这件事表面上只是“保存一下训练好的参数”实际上是对工程化能力的一次小考。把文件格式、结构注册、预处理配置、运行环境、备份策略这些细节都考虑周全你的模型才能从“只能在自己电脑上自嗨”真正变成“在任意环境里被人放心引用”的成果。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价