ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch CNN模型保存与加载:自动挑选最优权重实战指南

PyTorch CNN模型保存与加载:自动挑选最优权重实战指南 做CNN项目最难受的时刻是什么不是网络不收敛也不是显存不足而是模型明明训练得很好结果程序一崩、内存一清或者保存了一堆乱七八糟的权重文件最后根本分不清哪个才是最优模型。这一课我们就专门解决这个“最后一公里”问题在PyTorch里CNN的最优模型到底怎么保存、怎么加载、怎么在实战中自动挑出最好的一份留档。本文整个系列已经讲到第六课前几课我们把卷积、池化、激活函数、训练循环都过了一遍代码也能跑通了但很多人在模型持久化这个环节会卡住尤其是“早停时保存哪个epoch的权重”“加载报错怎么办”“继续训练时怎么恢复优化器状态”这类细节踩坑的人特别多。这篇文章我会从原理到代码把我自己实际用下来最顺手的方案完整拆开讲。1. 到底什么样的模型才算“最优模型”1.1 训练精度高不代表模型好用我在带新人做图像分类的时候经常看到有人盯着训练集准确率乐开花觉得自己模型已经完美了结果一到验证集上直接露馅。训练精度高只能说明模型学会了“背答案”不一定学得会“做新题”。真正的“最优模型”应该是在未见过的验证集上表现最好的那一份权重而不是训练集上最后一次迭代的权重。这里需要先达成一个共识我们评价CNN模型好坏看的是泛化能力。你可以用验证集准确率、验证集损失、F1分数等指标来衡量但核心原则是一样的——用验证集来当裁判让模型在训练过程中接受这个裁判的检验然后把检验成绩最好的那一份权重单独存下来。我自己的习惯是每个epoch结束之后用验证集跑一遍前向传播记录验证损失和验证准确率然后和当前的历史最佳成绩对比。如果更好了就覆盖保存一份“最优模型”如果没有变好就继续训练。这样不管最终训练到第100个epoch还是第200个epoch我手上永远握着的是验证集表现最好那一刻的权重。1.2 验证集、测试集要分清楚很多初学者会把验证集和测试集混着用这其实是个很危险的习惯。验证集是用来做模型选择和超参数调整的测试集是最终评估模型泛化能力的两者数据绝对不能交叉。我们选择“最优模型”时参考依据必须是验证集确定最优模型之后想对外宣称模型的效果再去跑一遍测试集。如果你数据集不大可以考虑用K折交叉验证来辅助判断但深度学习场景下通常按60%、20%、20%划分训练集、验证集、测试集就够用了。切数据的时候记得要做随机打乱并且固定好随机种子不然每次实验的“最优模型”可能都不一样复现起来非常头疼。1.3 最优模型不只是“准确率最高”这一个标准用验证集准确率作为指标是最直观的做法但并不是所有任务都适合只看准确率。做医学影像分类时阳性样本可能只占5%这时候准确率会很虚哪怕模型全部预测阴性也能拿到95%的准确率。这种场景下更好的做法是监控验证集上的F1分数或者AUC值把这些指标作为“最优模型”的评选标准。损失函数也可以参与评判。验证损失往往比准确率更敏感因为它能反映模型对预测结果的确信程度——准确率相同的情况下验证损失更低的模型通常预测得更有底气。我在实际项目中习惯同时记录验证准确率和验证损失保存的时候优先选择验证损失最小的权重因为准确率偶尔会出现小幅震荡验证损失的整体趋势更能反映模型的稳定状态。2. PyTorch模型保存的三种方式与适用场景2.1 只保存权重参数state_dict方案PyTorch里最推荐的保存方式是把模型的state_dict保存下来。state_dict本质上是Python的字典对象里面存了模型所有的可训练参数——也就是卷积核权重、偏置这些张量。一个CNN模型可能包含几百万甚至几千万个参数state_dict都会以键值对的形式存好。代码很简单torch.save(model.state_dict(), cnn_best.pth)加载的时候先重新创建模型结构再把权重文件里面的参数灌进去model CNNModel() model.load_state_dict(torch.load(cnn_best.pth)) model.eval()这个方案的优点是只存参数占用空间小灵活性高加载时模型结构由你的代码决定。缺点是你必须保证创建模型的代码和训练时的代码完全一致稍微改了一个卷积层的通道数加载就会报KeyError。我见过有人直接把整个训练脚本原封不动跑一遍最后在测试脚本里引用了另一份模型定义代码结果怎么都加载不上排查了半天发现是Relu写成了ReLU、类名虽然一样内部层名字对不上。所以用state_dict方案模型结构代码一定要统一管理。2.2 保存完整模型一键存档但坑也多PyTorch也支持直接保存整个模型对象代码是torch.save(model, cnn_full.pth)。加载的时候更简单一行代码model torch.load(cnn_full.pth)这种方法适合快速实验、临时保存不需要关心模型定义是否一致因为模型结构和参数都打包在一起了。但我不太建议在正式项目中用这个方案第一保存的文件体积明显更大因为模型结构、类定义引用全都序列化进去了第二加载时依赖的类路径必须和保存时一致你的模型类如果写在某个脚本里面之后脚本改了名字或者类挪了位置加载就会失败第三安全性差一些pickle反序列化存在被植入恶意代码的隐患虽然本地单机用问题不大但涉及分享权重文件时还是要谨慎。我自己一般只在Demo演示时用torch.save(model)这种写法正规训练流程一律只用state_dict或者checkpoint方案。2.3 保存Checkpoint让训练断点续传所谓checkpoint就是把整个训练状态打包保存下来。这里不光有模型的state_dict还有优化器的state_dict、当前的epoch数、验证集历史最佳成绩、随机数生成器的状态等等。这样做的最大价值是训练中途断电、服务器重启、显存溢出导致进程被杀的时候你可以从checkpoint继续训练而不是从头再来。一个标准的checkpoint字典长这样checkpoint { epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_val_acc: best_val_acc, best_val_loss: best_val_loss, scheduler_state_dict: scheduler.state_dict(), } torch.save(checkpoint, checkpoint_epoch50.pth)加载checkpoint恢复训练时代码要做反向操作checkpoint torch.load(checkpoint_epoch50.pth) model.load_state_dict(checkpoint[model_state_dict]) optimizer.load_state_dict(checkpoint[optimizer_state_dict]) start_epoch checkpoint[epoch] 1 best_val_acc checkpoint[best_val_acc]这里有个容易忽略的细节优化器的state_dict里记录了每个参数的动量、学习率缓存等信息。如果不恢复优化器状态就直接继续训练动量信息会丢失训练曲线会出现一段非常诡异的反弹损失先是飙升然后再慢慢降低。恢复学习率调度器也一样不然你的学习率会从初始值重新开始和当前训练阶段完全不匹配。2.4 三种方式的横向对比为了方便选择我把三种方式放在一张表里对比保存方式文件内容恢复难度适用场景state_dict模型权重参数需先重建模型结构推理部署、迁移学习、权重交换完整模型模型结构权重一行代码加载快速Demo、临时保存checkpoint模型优化器epoch指标需手动解包字典训练中断恢复、长期训练任务记住一个口诀训练中用checkpoint训练后留state_dict图省事用完整模型但不推荐。3. 自动保存最优模型的完整训练循环3.1 为什么要自动保存而不是手动保存很多人写训练代码时习惯在每个epoch结束之后保存一个“当前模型”文件名起成model1.pth、model2.pth……跑到最后盘里堆了几十个权重文件但让你挑一个效果最好的你根本说不清。还有人是等到训练全部结束再保存但最后一个epoch的权重往往不是验证集表现最好的早停的时候尤其如此。手动保存的问题在于你不知道哪个epoch会收敛到最佳状态也不可能每个epoch都守着看验证集指标。所以自动保存最优模型是工程上的必然选择。逻辑很简单每个epoch结束后在验证集上算一次指标如果比历史最佳成绩更好就覆盖保存。我喜欢用更稳的方式每次发现新的最佳成绩时保存一份独立的文件文件名里带上验证准确率比如best_model_val_92.31%.pth。这样就算后面过拟合了、模型继续训练变差了你手上依然留着一份历史最佳权重不会因为覆盖保存而丢失。3.2 训练循环里的关键代码这里给出一个我自己实际在用的完整训练循环框架以PyTorch为例。这个框架融合了自动保存最优模型、早停机制和学习率调度best_val_loss float(inf) best_val_acc 0.0 patience 10 early_stop_count 0 for epoch in range(1, num_epochs 1): # 训练阶段 model.train() train_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() train_loss loss.item() # 验证阶段 model.eval() val_loss 0.0 correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() val_loss val_loss / len(val_loader) val_acc 100.0 * correct / total # 自动保存最优模型 if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_model_acc{:.2f}.pth.format(val_acc)) early_stop_count 0 print(Epoch {}: 新的最优模型验证准确率 {:.2f}%.format(epoch, val_acc)) else: early_stop_count 1 print(Epoch {}: 验证准确率 {:.2f}%未提升.format(epoch, val_acc)) # 学习率调度 scheduler.step(val_loss) # 训练中断恢复用的checkpoint torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_val_acc: best_val_acc, scheduler_state_dict: scheduler.state_dict(), }, lastest_checkpoint.pth) # 早停 if early_stop_count patience: print(验证准确率连续{}个epoch未提升提前停止训练.format(patience)) break这个框架里有几个细节我想重点说明。第一个是model.train()和model.eval()的切换很多人漏掉这一句导致Dropout和BatchNorm在训练和验证时表现不一致验证指标忽高忽低模型明明不错却保存不下来。第二个是torch.no_grad()验证阶段不需要计算梯度加上它不仅能省显存还能大幅提升验证速度。第三个是scheduler.step(val_loss)很多学习率调度器依赖验证损失来调整学习率把验证损失传进去才能让调度器正常工作。3.3 早停和保存模型的协作逻辑上面代码里我放了两个机制自动保存最优模型和早停。为什么要配合使用因为两者解决的是不同的问题自动保存解决的是“最好的权重在哪”的问题早停解决的是“什么时候该停止训练”的问题。如果只做早停不自动保存你会在某个时刻停止训练但停止时的权重可能已经过了最佳状态一段距离如果只自动保存不早停你会继续训练很久浪费时间还容易过拟合。两个机制结合起来代码会在验证指标连续多个epoch不提升时停止训练但停止时手上拿到的仍然是历史最佳权重。对于early_stop_count的初始值这里有个小坑第一次epoch验证准确率如果没有比初始值0更好——虽然基本不可能但只要没提升计数就会加1。实际上第一个epoch之后验证精度必然远超0所以这个逻辑是安全的。如果你把best_val_acc初值设成100.0那模型永远不会保存这个细节值得注意。3.4 多指标联合评判时怎么保存前面说过有些任务只看准确率不够我这里再给一个更通用的方案保存时同时参考验证损失和验证准确率优先级是验证损失优先。下面的代码片段展示了这个思路if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), best_model_loss{:.4f}.pth.format(val_loss))这样训练过程中会保存两份文件一份是验证准确率最高的一份是验证损失最低的。最后测试的时候把两份都拿来跑一遍测试集选成绩更好的那一个稳。实际操作中虽然验证损失和验证准确率通常同步变化但偶尔会出现验证损失低但准确率也偏低的情况这时候需要你根据任务目标做取舍比如误判代价高的场景优先选损失低的模型均匀分类场景优先选准确率高的模型。4. 模型加载与使用的四大场景详解4.1 场景一加载最优模型做推理预测训练完之后你要用最优模型对新图片做预测。这个场景最关键的是加载之后要调用model.eval()。PyTorch的模型默认处于训练模式如果你加载之后直接推理BatchNorm会使用当前batch的统计量而不是训练集累积的统计量Dropout也还在随机失活输出结果会不稳定。一个完整的推理流程import torch from PIL import Image import torchvision.transforms as transforms device torch.device(cuda if torch.cuda.is_available() else cpu) model CNNModel() model.load_state_dict(torch.load(best_model_acc92.50.pth, map_locationdevice)) model.to(device) model.eval() transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) image Image.open(cat.jpg) input_tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): output model(input_tensor) predicted_class torch.argmax(output, dim1).item()这里的input_tensor是四维的形状是batch_size, channels, height, width即(1, 3, 224, 224)。为什么不能省略unsqueeze(0)因为模型在训练时就是按batch喂数据的卷积层和全连接层都期望输入带上batch维度直接喂3维数据会报错。4.2 场景二加载checkpoint继续训练继续训练是checkpoint方案的主场。加载checkpoint的时候有一点要特别留意保存的文件里已经包含了模型和优化器当前状态你不需要重新初始化它们但需要重新把模型切回训练模式checkpoint torch.load(lastest_checkpoint.pth, map_locationdevice) model.load_state_dict(checkpoint[model_state_dict]) optimizer.load_state_dict(checkpoint[optimizer_state_dict]) scheduler.load_state_dict(checkpoint[scheduler_state_dict]) start_epoch checkpoint[epoch] 1 best_val_acc checkpoint[best_val_acc] model.to(device) model.train() for epoch in range(start_epoch, num_epochs 1): # 正常训练流程 ...这里写着start_epoch checkpoint[epoch] 1是让你从保存的位置继续往下训练而不是重新从第0个epoch开始。如果你不恢复epoch计数后续输出日志会重复但更关键的是学习率调度器的进度会同epoch计数绑在一起错乱的epoch会让调度器行为异常。继续训练还容易遇到一个隐性bug保存checkpoint时模型在cuda上恢复时加载到了一个纯CPU的设备上。你需要在torch.load中传入map_location参数做张量定位否则会报“Attempting to deserialize object on a CUDA device”的错误。这个报错我用太多次了几乎每个跑过分布式训练的人都会遇到。4.3 场景三迁移学习与特征提取CNN项目里很常见的一个用法是把预训练好的模型权重加载到新模型上只替换最后的全连接层。比如你用ImageNet上训练好的ResNet来初始化一个猫狗分类器这时候只需要加载state_dict中除分类器以外的部分。代码实现pretrained_dict torch.load(resnet50_imagenet.pth, map_locationdevice) model_dict model.state_dict() pretrained_dict {k: v for k, v in pretrained_dict.items() if k in model_dict and model_dict[k].shape v.shape} model_dict.update(pretrained_dict) model.load_state_dict(model_dict)这种筛选式加载的好处是模型的分类器结构变了也不会报错只会加载共享的部分参数。shape不一致的层也会被过滤掉避免因为输出的类别数不同导致加载失败。这就是为什么用state_dict做迁移学习比保存完整模型更灵活。4.4 场景四导出到生产环境做部署如果你需要把CNN模型部署到服务端用Python脚本直接加载PyTorch模型当然可以但生产环境更常用的是ONNX或TorchScript格式。这两个格式不依赖原始模型类定义可以在不同的推理框架里运行。导出TorchScript的代码scripted_model torch.jit.script(model) scripted_model.save(model_jit.pt)导出ONNX的代码dummy_input torch.randn(1, 3, 224, 224, devicedevice) torch.onnx.export( model, dummy_input, model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )导出ONNX的时候一定要给一个形状正确的dummy_input模型会沿着dummy_input的维度做一次前向传播来梳理计算图。注意动态轴dynamic_axes的配置如果不配导出的模型固定batch为1线上请求来了没法一次推理多张图。这个细节在做服务化部署时很重要。5. 常见问题与排查技巧实录5.1 动不动就报Missing key或Unexpected key这个报错信息我几乎每周都会看到“Missing key(s) in state_dict: conv1.weight...”或者“Unexpected key(s) in state_dict: fc.weight”。出现这个问题的原因只有一个模型结构定义和保存时不一致。排查的思路从三个方向走第一确认你加载权重时创建的model类和你训练时用的类是不是同一个文件里的同一个类复制粘贴到新脚本时最容易出这种问题第二确认类的内部层名字是否一致比如把self.conv1改成了self.layer1第三确认输出类别数是否变了classification模型的fc层权重形状会随类别数变化。解决的办法是删掉模型缓存文件或者重新核对代码。有些人会在load_state_dict时加上strictFalse来强行加载但我不推荐在正式项目里这么干因为那会让部分参数保留随机初始化状态模型是“半残”的。5.2 加载模型时一直在报CUDA相关错误如果你在CPU机器上加载GPU训练的模型会遇到“RuntimeError: Attempting to deserialize object on a CUDA device”的报错。这是因为保存state_dict时记录了一个写入设备的信息默认情况下加载时会尝试把张量放回原来的设备。解决办法就是在torch.load中指定map_locationmodel.load_state_dict(torch.load(model.pth, map_locationtorch.device(cpu)))反过来如果你在GPU环境加载之前用CPU训练的模型问题不大PyTorch会自动帮你把模型挪到GPU前提是你加载之后对模型调用了model.to(device)方法。5.3 加载后推理准确率掉了很多这种问题十有八九是忘了调用model.eval()。模型默认是训练模式BatchNorm和Dropout的行为都和推理时不一样直接导致输出不稳定。我自己做图像分类时加载完模型的第一件事就是写model.eval()这个习惯已经刻进DNA了。还有一种情况是数据预处理对不上。测试时用的Normalize均值和标准差如果和训练时不一样模型对输入数据的分布认知就会错乱准确率暴跌。建议把数据预处理的参数统一封装成一个函数训练和推理都调用同一个函数。5.4 训练很久但最优模型始终没有覆盖如果自动保存的if判断一直不触发先检查best_val_acc的初始值。很多人把初始值设成了100.0导致第一轮验证准确率无论如何都比100.0小永远不满足覆盖条件。正确做法是把初始值设成0.0或者一个极小的负数。另外如果是使用验证损失做判断初始值要设成float(inf)不要设成0.0否则同样会导致永远不保存。5.5 训练中断后恢复的checkpoint无法继续训练我遇到过一个比较隐蔽的场景数据集加载器每次运行的时候都从头开始打乱顺序恢复训练虽然模型和优化器都恢复到了断点时的状态但你喂给模型的图片顺序和之前完全不同。这个问题不影响模型收敛因为数据始终是随机顺序但如果你追求完全复现原训练曲线建议在保存checkpoint时把DataLoader的随机状态一并保存。做法是在训练循环开头保存torch.random.get_rng_state()恢复时用torch.random.set_rng_state()复原。不过常规训练不追求这个可以按需处理。5.6 一张问题排查速查表现象可能原因优先排查方向Missing key报错模型结构不一致核对模型定义代码与训练时是否一致CUDA反序列化报错设备不匹配torch.load加map_location参数推理准确率暴跌忘了model.eval()加载后切换到评估模式最优模型从不更新best_val_acc初值错误检查初值是否设成0.0或-inf对应指标checkpoint恢复后损失飙升优化器状态没有恢复恢复optimizer和scheduler的state_dict保存文件特别大保存了完整模型换用state_dict方案加载后模型参数没变没有调用optimizer.zero_grad检查训练循环中梯度清零逻辑5.7 一个独家建议文件名里带指标很多人在一次训练中保存几十个模型之后根本分不清哪个对应哪个。我的习惯是把关键指标直接写进文件名比如best_model_acc92.50.pth、best_model_loss0.2314.pth、checkpoint_epoch50.pth。看起来文件名长了一点但之后翻目录、写实验记录、发ntu有经验的工程师都这么干谁用谁知道。6. 写在最后的实操体会这一课的内容看起来多但核心就三句最优模型用验证集选保存权重用state_dict训练和推理的状态切换别含糊。我自己刚开始做CNN的时候是用Excel手动记录每个epoch的验证精确率再手动复制模型文件效率低还容易出低级错误。后来把自动保存逻辑写进训练循环之后重新跑实验就舒服多了晚上开个训练任务第二天早上起来直接看best_model文件不用担心跑过头或者存错。最后再分享一个小技巧如果训练数据和验证数据的分布差异比较大自动保存的最优模型可能会在训练后期才被刷新。这种时候不要慌先看验证损失的曲线是不是还在震荡如果已经平稳下降、准确率不再提升就可以直接停止训练了已经保存的最优模型足够用。希望这篇内容能帮你在做CNN实战时少走一些弯路遇到模型保存和加载的问题回来翻翻这篇文章对照排查一遍就好。
返回列表