搞PyTorch这几年,被问得最多的一个问题就是“模型怎么保存怎么加载”。说实话,这话题看着简单,但真到了实战里,坑一点都不少。我印象最深的一次,是跑了一个两天的训练任务,结果第三天迭代代码的时候发现原脚本被覆盖,模型文件也没存对位置,整个结果直接没了。后来我养成一个习惯:不管项目多小,训练一启动就先写好保存逻辑,甚至先空跑一遍确认能存、能读、能恢复。就是这个“无聊”的习惯,帮我省下了不知道多少个重训的夜晚。
这篇文章不整那些虚的,直接围绕“Pytorch保存和加载模型”这个主题,把两种核心保存方式、checkpoint恢复、分布式训练的坑、模型文件体积优化、以及版本升级带来的兼容性问题全部梳理一遍。不管你是刚入门PyTorch的新手,还是写了不少训练脚本但没系统整理过保存策略的老手,这篇文章都能让你少踩几个坑,把时间花在真正有用的地方。
1. 先搞清楚两种保存方式,再动手写代码
很多教学帖一上来就给代码,但没解释为什么有两种保存方式,区别是什么。这一步搞不清楚,后面遇到问题就是两眼一抹黑。
1.1 state_dict:只存参数,不存结构
state_dict是PyTorch中一个非常核心的概念。简单说,它就是模型里所有可学习参数(weight、bias、BN层的running_mean等)组成的一个Python字典对象,key是参数名,value是Tensor。保存模型时,推荐的做法是只保存这个字典,而不是把整个模型对象都序列化。
为什么推荐只存state_dict?它的数据量小、结构清晰、加载灵活,而且和代码结构解耦。你训练时定义了一个五层卷积的模型,只要在加载时用相同的类定义重新实例化一个模型,然后调用load_state_dict,参数就能一一对上。这就好比你搬家时只搬一箱衣服(参数),而不是把整个衣柜(模型结构)都搬走。只要到了新家重新组装一个衣柜,衣服直接放进去就行。
保存state_dict的代码非常简单:
python复制# 保存
torch.save(model.state_dict(), 'model_weights.pth')
# 加载
model = MyModel() # 需要先实例化一个结构一致的模型
model.load_state_dict(torch.load('model_weights.pth'))
model.eval()
这段代码几乎出现在所有PyTorch教程里,但实际操作中会衍生出很多细节,比如模型类定义位置变了怎么办、实例化时参数不一致怎么办、加载后需不需要调model.eval()。这些细节我放在后面专门讲,先记住一个原则:state_dict保存方式的关键前提是“模型结构在加载时已知且一致”。
1.2 保存完整模型:简单但脆弱
另一种方式是直接保存整个模型对象:
python复制# 保存
torch.save(model, 'model_full.pth')
# 加载
model = torch.load('model_full.pth')
model.eval()
这种方式的优点是非常省事,不需要在加载时重新实例化模型,模型结构也跟着文件走。但代价是文件体积大(会把很多附属信息也存进去)、兼容性差(PyTorch版本升级后很可能加载失败)、而且容易把模型类定义里的临时状态一并序列化,导致加载后模型行为异常。
我自己的经验是,这种方式适合快速演示、或者模型结构本身非常简单且后续不会再改代码的情况。一旦进入正式项目和长期迭代,我强烈建议用state_dict方式,因为模型结构是代码的一部分,应该由代码管理,而不是被埋在序列化文件里。
1.3 一张表看清两者的取舍
| 对比维度 | state_dict | 完整模型 |
|---|---|---|
| 文件体积 | 较小,只含参数 | 较大,包含结构信息 |
| 加载灵活性 | 需先定义模型类 | 不需要 |
| 版本兼容性 | 较好 | 较差,跨PyTorch版本易挂 |
| 代码可维护性 | 好,结构由代码控制 | 差,结构被文件“锁死” |
| 推荐场景 | 正式项目、长期迭代 | 快速Demo、纯研究验证 |
注意:无论是哪种方式,保存之前最好先确认目标保存目录存在,否则会直接抛出
FileNotFoundError。通常是提前用os.makedirs(save_dir, exist_ok=True)处理一下,这点小细节天天有人踩。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 保存模型的标准姿势与工程化细节
知道两种方式之后,接下来要做的是把保存动作做成一套规范流程,而不是训练结束随手保存一下。工程化的保存策略能让你在断点恢复、调参回溯、多轮实验对比时不被动。
2.1 基础保存代码与目录规划
先看一个工程化一点的保存写法:
python复制import torch
import os
class Trainer:
def __init__(self, model, save_dir="./checkpoints"):
self.model = model
self.save_dir = save_dir
os.makedirs(save_dir, exist_ok=True)
def save_model(self, epoch, filename=None):
if filename is None:
filename = f"epoch_{epoch}.pth"
save_path = os.path.join(self.save_dir, filename)
torch.save(self.model.state_dict(), save_path)
print(f"Model saved to {save_path}")
这段代码本身不复杂,但有几个值得注意的点:
第一个,为什么保存目录要单独抽象出来。如果你直接写死torch.save(model.state_dict(), 'model.pth'),后面做实验对比时会非常痛苦。几个模型文件全堆在工作目录里,名字要么是model.pth要么是model_final.pth,根本分不清哪个对应哪次实验。我在实际项目中习惯用“项目名_日期_epoch数”的结构,比如20241005_resnet50_epoch_20.pth,这样文件本身就有信息量。
第二个,保存路径里尽量不要出现中文和空格。虽然Linux和Windows现代文件系统都支持,但有些工具链和后续脚本处理起来会出幺蛾子。长期做模型管理,路径命名越保守越省心。
第三个,保存之前先跑一次加载流程做验证。这一点可能很多人觉得没必要,但我的建议是,在模型刚初始化时先保存一次、再加载一次、确认能跑通,再做正式训练。特别是当你模型里用了自定义层、动态结构或者forward里有临时变量时,提前验证能避免训练一天后发现根本加载不回来。
2.2 训练中断点(checkpoint)的科学保存
如果你只是保存模型参数,那训练中断后想恢复会有个大麻烦:优化器的状态丢了。Adam优化器里有动量(momentum)和二阶矩估计,这些信息在训练中途被丢弃,恢复训练时会有一段时间的“冷启动”,效果明显变差。
科学的做法是把训练状态打成一个“断点包”(checkpoint),里面至少包含以下内容:
python复制checkpoint = {
"epoch": epoch,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"best_acc": best_acc,
"lr": scheduler.get_last_lr() if scheduler else None,
"config": {"model_name": "resnet50", "input_size": 224}
}
torch.save(checkpoint, f"checkpoint_epoch_{epoch}.pth")
这里面最关键的是optimizer_state_dict。我之前有一次图省事,只保存了模型参数,训练中断后直接从epoch 10继续跑,结果后面几轮的loss波动非常大,因为Adam里累积的梯度信息全部重置了,等于带着新优化器从半路开始。所以只要你有“恢复训练”的需求,优化器状态必须一起存。
scheduler的学习率状态也应该一起存,否则恢复训练时学习率调度会从头开始算,实际学习率可能和你预期的完全不一样。特别是用ReduceLROnPlateau这类动态调整策略时,不保存scheduler状态基本等于调参白做。
2.3 恢复训练时,一个都不能少
加载checkpoint继续训练,不是简单model.load_state_dict就完事了,你需要把整个训练状态都恢复回来:
python复制def resume_training(checkpoint_path, model, optimizer, scheduler):
checkpoint = torch.load(checkpoint_path)
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
if scheduler and "lr" in checkpoint:
scheduler.load_state_dict(checkpoint["lr"])
start_epoch = checkpoint["epoch"] + 1
best_acc = checkpoint["best_acc"]
return model, optimizer, scheduler, start_epoch, best_acc
这里有个小细节值得注意:如果你用的是CosineAnnealingLR这类需要知道总训练轮数的调度器,恢复训练时传入的T_max必须和保存时保持一致,否则学习率曲线就乱了。这种问题排查起来很隐蔽,因为你只看loss曲线会发现它“看起来正常”,但收敛速度已经变了。
提示:写checkpoint时,我习惯在dict里加一个
"config"字段,记录模型类型、输入维度、关键超参等元信息。加载时先检查config是否和当前模型定义匹配,不匹配就直接报错。这个习惯在团队协作和模型交接时特别有用,不然拿到一个.pth文件根本不知道它原来是什么结构。
3. 加载模型的完整姿势与易错点
加载模型表面上是保存的逆操作,但实际操作中踩坑概率比保存高得多,尤其是在模型定义、设备管理和版本兼容这几个环节。
3.1 加载state_dict的标准流程
加载state_dict的标准流程分成三步:
第一步,实例化一个“结构一致”的模型。这里要注意,PyTorch的load_state_dict只负责把参数值填进去,它不会检查你的模型输入输出形状是否正确。如果你实例化模型时把num_classes从1000改成了10,只要参数名对得上、维度刚好一致,load_state_dict不会报错,但模型逻辑可能已经错了,这种错比报错更难发现。
第二步,调用load_state_dict。建议使用严格模式(默认就是strict=True)。如果保存和加载的状态字典中有key对不上,PyTorch会明确告诉你缺了哪些、多了哪些,这对排查问题非常有帮助。千万不要随手改成strict=False去掩盖问题,除非你明确知道自己在干什么。
第三步,根据使用场景调用model.eval()或model.train()。这一步是无数新手翻车的地方。model.eval()主要影响Dropout和BatchNorm等层的行为——Dropout会关闭随机失活,BatchNorm会使用累计的running_mean和running_var。如果你加载模型是为了做推理和评估,不调用eval()的话,结果可能和训练时的验证指标完全不同,尤其是有BatchNorm的模型,差距特别明显。
python复制# 完整标准流程
model = ResNet50(num_classes=1000)
state_dict = torch.load("model_weights.pth", map_location="cpu")
model.load_state_dict(state_dict)
model.eval() # 做推理时必须加上
3.2 map_location:解决设备不一致的万能钥匙
在保存模型时,模型参数在哪个设备上(GPU/CPU)是会被记录进文件里的。如果保存时在GPU上,加载的机器没有对应GPU、或者GPU序号不一样,直接torch.load(path)大概率报错。
map_location参数就是用来处理这个问题的:
python复制# 从GPU保存的文件,在只有CPU的机器上加载
model.load_state_dict(torch.load("gpu_weights.pth", map_location="cpu"))
# 在另一张GPU上加载,指明设备
model.load_state_dict(torch.load("gpu_weights.pth", map_location="cuda:0"))
我在实际项目里的习惯是,统一先加载到CPU,再显式调用.to(device)。这样不管是训练还是推理,流程都干净可控,不会出现加载后模型在GPU上、但后续流程默认它在CPU上的错位问题:
python复制device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = ResNet50()
model.load_state_dict(torch.load("weights.pth", map_location="cpu"))
model.to(device)
3.3 weights_only:PyTorch版本升级带来的新参数
如果你用的是新版PyTorch(2.x),torch.load会出现一个新参数weights_only。这个参数在PyTorch 2.6之后默认变成了True,目的是防止加载恶意pickle文件导致代码执行。但副作用是,如果你的checkpoint里除了权重还存了optimizer.state_dict、自定义类、甚至torch.device对象,直接加载会报错。
如果你加载的是老的checkpoint(包含自定义类或完整对象),需要显式指定weights_only=False:
python复制checkpoint = torch.load("checkpoint_epoch_10.pth", map_location="cpu", weights_only=False)
不过安全起见,我更推荐的做法是:尽量让checkpoint里只存基础类型和Tensor,自定义类不要直接塞进去。这不仅是安全问题,也是兼容性问题。举个我踩过的坑:我在模型类里定义了一个lambda函数作为激活函数,结果保存checkpoint后,修改了类的代码再加载,pickle怎么都找不到原来的lambda引用,直接报错。从那以后,我在checkpoint里只保存参数和基础数据类型,需要什么配置用config字典单独记录。
4. 训练实战:保存加载如何融入日常流程
保存和加载不是两个孤立动作,它们应该嵌入训练和评估的完整流程。这一节我把在实际训练中怎么用这两件事讲清楚。
4.1 边训练边保存与最佳模型筛选
训练过程中,模型在不同epoch的表现是波动的。通常我不会等训练结束才保存最终模型,而是每个epoch或每几个epoch保存一次checkpoint,同时维护一个“最佳模型”的单独备份。
一个常见的实现思路是:
- 每个epoch结束后计算验证集指标(如accuracy/mAP等)。
- 如果当前指标优于历史最佳,就覆盖保存
best_model.pth。 - 每隔固定epoch数保存一个带时间戳的断点,用于中途恢复。
python复制best_acc = 0.0
save_dir = "./checkpoints"
for epoch in range(start_epoch, total_epochs):
train_one_epoch(...)
val_acc = validate(...)
if val_acc > best_acc:
best_acc = val_acc
torch.save(model.state_dict(), os.path.join(save_dir, "best_model.pth"))
if epoch % 5 == 0:
torch.save({
"epoch": epoch,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"best_acc": best_acc,
}, os.path.join(save_dir, f"checkpoint_epoch_{epoch}.pth"))
这样的好处是:即使训练中断,你也有最近的断点可以恢复;而best_model.pth永远指向验证集上表现最好的那版参数,用于后续测试和部署。这里有一个实操心得:当训练快结束时,我会特意把最后一个epoch的checkpoint也保留下来。有时候验证集上“最好”的模型不一定泛化最好,最后一个epoch的模型可能在某些场景下更稳。
4.2 加载模型做评估时,确保“所见即所得”
我见过不少同学加载模型做评估时,指标比保存时低很多,然后怀疑是保存过程出了问题。其实大概率是评估流程的预处理和训练时不一致,或者是没有做model.eval()。
有一个小技巧很实用:加载模型后,先对同一条数据进行一次前向推理,看输出和保存前是否一致。如果你的模型在eval模式下是确定性的(没有Dropout和随机采样),两次输出应该完全一致。如果有差异,就说明评估流程中某个环节有随机性,或者模型模式没设置对。
4.3 数据并行/分布式训练下的保存与加载细节
在单卡环境下,模型参数名就是模块名,比如conv1.weight、fc.bias。但在DataParallel或DistributedDataParallel(DDP)环境下,模型会被包一层module,参数名就变成了module.conv1.weight。
如果你用DataParallel训练,保存时直接存model.state_dict(),那保存下来的key是module.conv1.weight。换到单卡环境加载时,load_state_dict会报“missing keys”和“unexpected keys”,看起来非常迷惑。
解决办法是保存时剥离module前缀,或者加载时做处理。比较推荐的做法是在保存前都统一存原始模型:
python复制# DataParallel场景
model = nn.DataParallel(model)
# 训练...
# 保存时取原始模型的state_dict
torch.save(model.module.state_dict(), "weights.pth")
DDP场景下,保存更稳妥的方式是存模型的原始module参数。因为DDP只是包装器,真正的参数都在model.module里。而加载时如果你用的是单卡代码,直接model.load_state_dict(torch.load("weights.pth"))就能对上。
另外,DDP训练时的optimizer最好也是在model.module上创建,否则参数引用会混乱。这个和保存加载直接相关,因为优化器状态里的参数索引对不上,恢复训练时会出现“Size mismatch”。
我建议在项目里封装一个统一的保存函数,自动判断是否需要剥离module前缀:
python复制def save_model(model, path):
raw_model = model.module if hasattr(model, "module") else model
torch.save(raw_model.state_dict(), path)
这样不管用的是单卡还是DDP,保存出来的文件都是干净的、不带module前缀的state_dict,加载逻辑就能保持简单统一。
5. 常见问题与排查技巧实录
保存和加载模型报错时,错误信息往往会直接给出大量“key”相关的细节,很多人看到这种错误就懵了。其实这些错误信息恰恰是排查的关键,学会读懂它们,你能省下大量时间。
5.1 最容易翻车的错误与修复方案
下面这份速查表,基本覆盖了我日常被问到的90%的问题:
| 错误现象 | 根本原因 | 解决方案 |
|---|---|---|
Missing key(s) in state_dict |
模型结构和保存时不一致 | 检查num_classes、层数等定义是否变化;确认是从原始模型保存的 |
Unexpected key(s) in state_dict |
加载了带module.前缀的参数,或混入了额外参数 |
检查是否用了DataParallel/DDP保存;用剥离前缀的方式处理 |
Can't get attribute 'xxx' on module |
checkpoint里存了自定义类或lambda函数 | 单独保存参数和config;加载时指定weights_only=False |
RuntimeError: Attempting to deserialize object on a CUDA device |
保存时在GPU,加载机器没有GPU或设备号不对 | 加载时加map_location="cpu" |
CUDA out of memory(加载后推理) |
模型被加载到显存,但显存不足 | 尝试加载到CPU推理,或减小batch size、用half精度 |
| 加载后指标和保存不一致 | 未调model.eval(),或预处理不一致 |
确认推理模式下调用eval();对比数据预处理流程 |
5.2 模型文件体积与加载速度优化
如果你训练的是大模型(比如超过1GB的模型),torch.save和torch.load会明显变慢,磁盘占用也大。这里有几个优化方向:
第一,只保存必要内容。很多人习惯把验证集预测结果、日志等塞进checkpoint,一个checkpoint存出好几个GB。实际上这些信息完全可以在训练时单独落盘,checkpoint里只留模型参数、优化器状态、epoch和指标就够了。
第二,用半精度保存。如果模型参数是FP32的,你可以转成FP16保存,体积直接减半。加载后再根据应用场景决定转回FP32还是直接用FP16推理:
python复制# 保存时转半精度
fp16_state_dict = {k: v.half() for k, v in model.state_dict().items()}
torch.save(fp16_state_dict, "model_fp16.pth")
但这招不是所有场景都适用。半精度保存有精度损失,尤其对某些对数值敏感的训练任务,加载后会有一点偏差。如果只是做推理展示,问题不大;如果要继续训练,我建议还是保留FP32版本。
第三,考虑用torch.save时传入_use_new_zipfile_serialization=True(默认就是),这是PyTorch 1.5之后的新格式,体积更小、加载更快。如果你加载老文件时遇到格式兼容问题,可以考虑用脚本转换一次。
5.3 版本兼容与跨环境迁移要点
PyTorch的存储格式并不保证向前兼容,这意味着用2.0保存的模型文件,在1.8上可能加载不了。遇到这种问题,有几个实用技巧:
一是尽量锁定PyTorch版本。在训练项目的requirements.txt里写明torch==2.1.0这类精确版本,团队内部统一。模型交接时,把PyTorch版本也一起告诉对方,能省掉一堆排查时间。
二是如果从老版本向新版本迁移,可以先在旧环境里把checkpoint加载出来,只保存state_dict(不存优化器状态),再用新环境加载。这样能规避不少pickle序列化带来的兼容性问题。
三是处理自定义算子。如果你用了torch.compile或自定义的CUDA扩展,保存和加载时可能需要重新编译或指定正确的扩展库。这种场景下,最稳的办法是保存一份纯参数版本用于交换,等目标环境配好了再加载。
5.4 一个卡了很久的加载报错案例
最后分享一个真实案例。有次我在服务器上训练了一个ResNet模型,保存后想下载到本地笔记本做演示。笔记本CPU环境加载时,一跑model.load_state_dict就报size mismatch for fc.weight,报错信息里还写着shape从(1000, 2048)变成了(10, 2048)。
排查了半天,发现是加载脚本里实例化模型时用了自己的默认参数num_classes=10,而训练脚本里用的是num_classes=1000。load_state_dict在参数维度不匹配时不会蒙混过关,会直接抛错,这个报错其实帮了大忙。但如果你把strict设成False,这个错误就会被吞掉,模型后半部分参数全是随机值,推理结果完全不可用。
这件事的教训是:定义模型时必须把关键超参统一管理,最好放进一个配置文件或常量里,不要散落在各训练脚本的__init__参数中。我现在的做法是每个项目都有一个config.py,模型初始化参数、路径、超参全在里面,保存加载时通过config里的信息重新实例化模型,从根上消除不一致问题。
结尾
保存和加载模型看起来是很小的知识点,但它是所有PyTorch项目的基石。如果你现在正被“模型保存后加载报错”折磨,我建议你先把state_dict和完整模型的区别搞清楚,再把checkpoint里的内容列一个清单,最后写一个统一的保存加载模块。这几件事做完,后续所有模型的复用、迁移和部署都会顺畅很多。
我个人实际操作中还有一个习惯,就是在第一次训练前先做个“保存-加载-评估”的闭环测试:保存一个刚初始化的模型,加载回来,对同一条数据跑预测,对比两次输出是否一致。这个习惯帮我抓到了很多早期bug,也推荐给你试试。模型训练是个漫长的过程,别让最后一步的保存加载毁掉整个项目。
