1. 保存和加载模型,为什么值得单独写一篇
先说个我见过太多次的场景:一台服务器上跑了十几个小时的训练,loss 从 2.3 慢慢降到 0.18,眼看着再有两个 epoch 就能收敛到理想水平了。结果第二天过来,发现终端被 systemd 重启踢了,训练进程直接没了。这时候你打开代码目录,发现没有任何 .pt 或 .pth 文件——因为训练脚本压根没写保存逻辑。崩溃不?
我之前带过不少做深度学习的实习生,发现一个很有意思的规律:大家训练模型的时候热情高涨,数据加载、网络结构、优化器调参都能折腾明白,但是一到"怎么把模型存下来、下次怎么加载"这步,就默认了torch.save(model, 'model.pth')然后torch.load完事。等真正部署、断点续训、跨机器迁移的时候,各种报错就冒出来了。
PyTorch 的保存和加载远不止"存一下读一下"那么简单,里面涉及了 state_dict 的设计哲学、张量的存储布局、设备之间的搬运、以及 checkpoint 里到底该放什么字段。这篇内容不绕弯子,直接把我这几年的实战经验整理出来,从最基础的 API 用法到各种踩坑现场,一次性讲清楚。无论你是刚入门 PyTorch 的小白,还是已经在做模型部署和迁移的工程师,这篇都应该能给你一些参考。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. PyTorch 保存加载的第一性原则:model 不是 Python 对象
2.1 你保存的不是模型,是模型的"参数快照"
很多框架老手第一次用 PyTorch 的时候,会很自然地把模型"整个保存"——就是torch.save(model, 'model.pth')。这用着确实方便,加载也简单,model = torch.load('model.pth')一行就回来了。但我要跟你认真聊聊这个做法背后藏着的问题。
PyTorch 官方推荐的做法是保存model.state_dict()。为什么?因为 state_dict 本质上是 collections.OrderedDict,里面存的可不是什么复杂的 Python 实例,而是参数名到张量的映射。这个设计有一个很务实的考虑:它把模型的网络结构(类定义、层顺序、forward 逻辑)和参数数据(权重、偏置、buffer)分开了。你可以把 state_dict 理解成"模型的配置数据",它脱离了网络结构本身,是一份纯粹的状态快照。
这里有个容易被忽略的技术细节:state_dict 里不止包含 nn.Parameter(就是那些 requires_grad=True 的权重和偏置),还包括模型的 buffer。最典型的就是 BatchNorm 层的 running_mean 和 running_var,这两个东西虽然不是梯度更新的目标,但在推理阶段对结果有决定性影响。这也是很多新手踩坑的地方——有人只手工保存了几个 torch.save 的张量,结果模型在训练时表现正常,一切换到 eval 模式做推理,输出结果完全不对,就是因为 BN 的统计量丢了。
2.2 为什么序列化 model 对象是个"定时炸弹"
torch.save(model, 'model.pth') 表面上能用,是因为 PyTorch 的 torch.save 底层用的是 Python 的 pickle 序列化。pickle 会把整个模型实例序列化成二进制。这带来一个特别难受的绑定:加载时必须保证代码环境里存在一模一样的模型类定义,且类的路径一致。
举个例子,你写了一个自定义模型类 TCNTransformer 放在 models/ts_model.py 里,训练完用 torch.save(model, 'model.pth') 存了下来。改天你想在另一个项目里加载这个模型,如果那个项目的代码里没有这个类,或者类定义路径变了(比如挪到了 models/backbone.py),加载时就会直接抛出 AttributeError: Can't get attribute 'TCNTransformer'。
还有版本兼容问题。PyTorch 本身迭代非常快,不同版本的源码里很多内部类的位置和结构都有变化。你用 1.13 训练的模型,隔半年用 2.x 的版本去 torch.load,报错概率不低。而 state_dict 就稳定得多,因为它只依赖张量字典结构和参数名,跨小版本加载基本没有任何问题。跨大版本只要没有破坏性的命名变化,也基本兼容。
所以在团队协作或者长期项目中,我的铁律是:存储用 state_dict,传输也只用 state_dict。模型结构定义是代码的事,代码有 Git 管,状态才需要文件管。
2.3 从 state_dict 的角度理解"模型结构必须一致"
加载 state_dict 的时候,PyTorch 是按参数名匹配的,不是按位置。这意味着模型的类定义可以先创建,加载只是填充数值。
举个例子,你定义一个网络:
python复制import torch
import torch.nn as nn
class MyNet(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(128, 64)
self.fc2 = nn.Linear(64, 10)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
model = MyNet()
state = torch.load('model.pth', map_location='cpu')
model.load_state_dict(state)
这里 load_state_dict 做的,就是把文件里的 fc1.weight、fc1.bias、fc2.weight、fc2.bias 分别填入当前 model 实例对应的参数里。如果当前实例的层名不匹配,就会报错告诉你 Missing key(s) 或者 Unexpected key(s)。
这个机制理解透了,后面很多"迷惑行为"就都有了解释:为什么冻结部分层也能加载?因为加载是针对指定 key 的,不存在的 key 不会强制创建。为什么 fine-tune 时要把某个层的名字改掉(比如 fc2 改成 fc2_new)?因为只要名字不对,加载时就不会覆盖它,那层就自然从随机初始化开始训练了。
3. 保存的三种姿势,分别用于不同的场景
我开始做深度学习的前两年,保存模型基本靠运气。后来有一个项目需要给客户出交付件,要求模型文件够小、加载够快、跨环境不出幺蛾子,这才把 PyTorch 的保存姿势研究透了。总结下来,按使用场景分成三种。
3.1 推理部署:只存 state_dict
这是最推荐、也是官方最支持的姿势。训练完成后,把模型的参数快照存下来:
python复制torch.save(model.state_dict(), 'resnet50_cifar10.pt')
加载的时候,先根据代码定义实例化模型,然后把参数填进去:
python复制model = MyNet()
state_dict = torch.load('resnet50_cifar10.pt', map_location='cpu')
model.load_state_dict(state_dict)
model.eval()
这里有个细节:加载完参数后,要把 model 切到 eval 模式。这在有 Dropout 和 BatchNorm 的网络里尤其重要。如果你忘记 model.eval(),模型默认处在 training 模式,Dropout 会随机丢神经元,BatchNorm 会继续更新 running stats,推理结果就会变得不稳定。这是我在上线时踩过的最经典的一个坑,后来我干脆写了个函数封装加载+eval,一步到位。
3.2 完整保存:场景有限,但确实有它的用途
虽然官方不推荐频繁使用,但某些场景下 torch.save(model) 仍然是最快且最省事的选择。比如你只是在本地做探索性实验,同一个脚本里训练完了立刻加载继续验证,模型类定义必然存在,不会有路径问题。
但是遇到下面这些场景,强烈不建议完整保存:
- 项目长期迭代,代码版本不断变更
- 需要把模型文件发给他人在不同环境下使用
- 模型定义存放在 Jupyter Notebook 或临时脚本里,没进版本控制
- 甲方要求的交付物,几个月后再要你重新加载,类文件找不到或改过
另外一个隐藏问题:torch.save(model) 出来的文件通常比只存 state_dict 大不少。因为 pickle 把模型类定义也存进去了,包括一些冗余的元信息。我在一次交付时对比过,同一个 PyTorch 模型,完整保存 86MB,只存 state_dict 是 78MB,虽然差距不算离谱,但如果是上传云盘或者通过邮件传输,这个差异就比较扎眼。
3.3 断点续训:保存的不是一个文件,而是一个"现场"
第三种,也是最值得我们花心思的:训练中断恢复。很多新手存 checkpoint 的时候只保存模型的权重,恢复训练时发现"loss 怎么降得跟第一次训练差不多"——因为优化器的状态丢了,学习率调度器也从头开始了。
一个完整的 checkpoint,至少要包含以下字段:
| 字段 | 作用 |
|---|---|
model_state_dict |
模型当前参数 |
optimizer_state_dict |
优化器动量和梯度历史 |
epoch |
当前训练轮数 |
best_loss 或 best_metric |
用于保存最佳模型 |
scheduler_state_dict |
学习率调度器的内部状态 |
写出来大致像这样:
python复制checkpoint = {
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scheduler_state_dict': scheduler.state_dict(),
'best_loss': best_loss,
}
torch.save(checkpoint, f'checkpoint_epoch_{epoch}.pt')
恢复训练时:
python复制model = create_model()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100)
checkpoint = torch.load('checkpoint_epoch_50.pt', map_location='cpu')
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
这里有一个很反直觉的细节:恢复训练时需要重新创建优化器,再进行 load_state_dict。这个顺序不能反。你如果先 load optimizer 再 load 模型没问题,但 optimizer 内部有的参数(比如 param_groups)在实例化时依赖于模型参数的结构,所以必须先创建完整的模型、再创建优化器、再往里面填状态。
3.4 保存 Checkpoint 的过程性策略:别只存一份
很多人的习惯是每个 epoch 结束 torch.save 一次,路径固定不变。这其实是个大隐患。训练到第 80 个 epoch 时,第 10 个 epoch 的 checkpoint 已经被覆盖了。如果第 85 个 epoch 模型崩了(比如 loss 出现 NaN),你想回退到 50 甚至 30 的状态做诊断,根本没得选。
我的建议是采用"滚动保存 + 多个副本"的策略:
- 每 5 个 epoch 保存一个带 epoch 编号的 checkpoint
- 同时维护一个
best_model.pt,只有在验证集指标创新高时才覆盖 - 再加上一个
last_model.pt,每次都更新,防止程序中断时没有最新现场
这样实现逻辑在现在这个存储成本几乎不是问题的时代,能省掉非常多不确定的麻烦。我在和很多做时间序列预测的同事合作时,他们用 TCN、Transformer 这类网络做股票预测,一个模型经常要训几百个 epoch,如果断点续训没做好,前面十几小时的算力就白搭了。
4. 加载模型时的设备问题:GPU 和 CPU 的来回搬运
4.1 map_location 是干什么的
PyTorch 的 torch.load 有一个非常关键但是老被忽略的参数:map_location。
如果你在 GPU 上训练,保存的是 GPU 上的张量,保存时会拷贝到 CPU 再序列化到磁盘。加载的时候,默认行为是什么?直接加载到之前保存时的设备。也就是说,你在一张 4 卡 GPU 机上用 cuda:0 训练的模型,如果直接 torch.load,它大概率会被加载进 cuda:0,哪怕你当前环境只剩着一张卡、编号是 cuda:1,或者你压根是在 CPU 机器上做推理——这就会直接报错。
正确做法是显式指定:
python复制# 在只有 CPU 的机器上加载
model.load_state_dict(torch.load('model.pt', map_location='cpu'))
# 在有一张 GPU 的机器上加载
model.load_state_dict(torch.load('model.pt', map_location='cuda:0'))
map_location 的取值可以是字符串,也可以是一个函数。字符串当然最简单。函数可以实现更复杂的映射,比如多卡训练时保存的模型键名带 module. 前缀,你又想在单卡环境加载,可以写个函数去掉前缀。这些后面还会细说。
4.2 加载后 "cuda:0 不存在" 的报错现场
很多人会碰到过这个异常:
code复制RuntimeError: Attempting to deserialize object on a CUDA device but torch.cuda.is_available() is False.
翻译一下:文件里保存的 tensor 说明它之前活在 cuda 设备上,但当前环境根本检测不到 GPU。这基本就是两种原因:
- 加载时没传
map_location='cpu',而当前机器没有 GPU。 - 传了
map_location='cpu',但模型里某些 buffer 类型是 CUDA tensor,需要额外转换。
第 2 种情况比想象中常见。比如你自己写了一个模型,提前把某些常量注册为 buffer,且句柄是在 GPU 上创建的。map_location='cpu' 只能处理 torch.load 反序列化动作,它能把张量文件从 GPU 反序列化到 CPU。但是如果模型内部做了 buffer.to(device) 这类手动搬运,设备冲突还是会存在。这个场景我建议直接把模型类里的 device 管理和加载分开,不要在实例化时绑定设备。
4.3 先把模型放 CPU 再搬设备,省得内存爆炸
还有一个实操技巧:无论目标设备是什么,torch.load 都是先加载到 CPU 再搬到目标设备,这个行为是 PyTorch 的默认设计。所以你不需要担心先把模型加载到 CPU 再移到 GPU 会有什么额外性能损失。
我的推荐写法是:
python复制device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
state_dict = torch.load('model.pt', map_location='cpu')
model = create_model().to(device)
model.load_state_dict(state_dict)
先把文件里的 tensor 全部落到 CPU,构建模型,搬到目标设备,最后再加载参数。这样做的好处是:如果你目标设备是 GPU,加载过程不会因为在反序列化时直接创建 CUDA tensor 而产生不必要的显存占用,先创建到 CPU 再整体搬移,反而更可控。
我刚入行时踩过一个相关的坑:某次在 8 卡机器上训练,checkpoint 是在 cuda:0 上保存的,然后我换机器只带了 4 卡,GPU 编号不是从 0 开始,结果一加载就 OOM,原因是反序列化时准备把一个大 tensor 塞进 cuda:0,但那个卡上还跑着别的任务。后来所有加载都统一 map_location='cpu',再手动指定设备,问题就再也没出现过。
5. 多卡训练与分布式场景:那些带 module. 前缀的坑
5.1 DataParallel 保存的模型,键名带 module. 前缀
经常有人在 PyTorch 论坛问:训练的时候用了 nn.DataParallel,为什么保存下来的权重键名全带 module. 前缀?比如 module.fc.weight,加载到一个没包 DataParallel 的裸模型时直接报 Missing key(s)。
原因很简单:nn.DataParallel 是一个包装类,它内部的 module 属性才是真正的模型。DataParallel 自身没有 fc 这个子模块,fc 在它的子模块 module.fc 上。所以当你对 DataParallel 包装后的模型调用 state_dict(),所有键名都带上了 module. 前缀。
解决方案有几种:
方案一:保存的时候剥掉包装
python复制if isinstance(model, nn.DataParallel):
torch.save(model.module.state_dict(), 'model.pt')
else:
torch.save(model.state_dict(), 'model.pt')
这是最干净的做法,保存的是一份没有 module. 前缀的标准权重文件。
方案二:加载时用 strict=False 或者改键名
如果文件已经带着 module. 前缀保存了,加载时可以用一个 helper 函数去掉前缀:
python复制state_dict = torch.load('model.pt', map_location='cpu')
new_state_dict = {}
for k, v in state_dict.items():
new_key = k.replace('module.', '')
new_state_dict[new_key] = v
model.load_state_dict(new_state_dict)
不过在动手改键名之前,你可以先试一下 load_state_dict(state_dict, strict=False)。strict=False 的含义是,加载时忽略当前模型缺失的键和不匹配的键。这样写的好处是至少不会报错,但代价是模型某些层的参数没被加载,仍然是随机初始化。所以我建议:strict=False 只用于调试,不要在生产代码里无脑用。
5.2 分布式训练(DDP)保存与加载的正确体位
torch.nn.parallel.DistributedDataParallel(DDP)和 DataParallel 的逻辑一致,它也有 module 包裹问题。但 DDP 的推荐做法是在训练脚本中做一个统一判断:保存之前先把模型拉回 CPU,然后取下 DDP 包装,保存原生模型的 state_dict。
python复制if dist.is_initialized():
torch.distributed.barrier()
model = model.module
checkpoint = {'model_state_dict': model.state_dict(), 'epoch': epoch}
torch.save(checkpoint, f'epoch_{epoch}.pt')
加载时分成两步:第一步构造原生模型(不带 DDP),加载 state_dict;第二步再基于原生模型包装 DDP。顺序不能反,否则 DDP 会在初始化时给每个参数打上用于梯度同步的 hook,如果你先 DDP 再 load_state_dict,可能会导致不同卡上的模型初始化状态不一致。
5.3 冻结部分模型参数的加载技巧
热词里有人提到"pytorch冻结部分模型",这和保存加载关系其实很大。微调场景下,你要加载一个预训练模型,但不想更新某些层——比如用 ResNet 做图像分类时,backbone 参数保持不动,只训练最后的全连接层。
经典做法是:
python复制for name, param in model.named_parameters():
if name.startswith('fc'):
param.requires_grad = True
else:
param.requires_grad = False
但有个前提:你仍然必须先通过 load_state_dict 把预训练权重加载进来,冻结只是不更新,不是不加载。顺序是:先实例化模型 -> 再 load_state_dict -> 再冻结 -> 再创建优化器。如果你先把参数冻结了,然后 load_state_dict,参数仍然会被填充,但因为 requires_grad=False,weight decay 和更新都跳过它们,这个顺序上没有问题。
另一个更细的点:冻结后创建优化器时,只把 requires_grad=True 的参数传进去,这样不仅节省显存,还能避免意外更新。加载完以后过滤一遍:
python复制trainable_params = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.Adam(trainable_params, lr=1e-4)
如果优化器在冻结之前就创建好了,参数列表中如果包含了 requires_grad=False 的参数,Adam 可能不更新它们,但 weight decay 在某些实现里仍然会作用于所有参数,导致问题,所以顺序至关重要。
6. 跨版本加载、参数名冲突与模型升级的兼容方案
6.1 PyTorch 版本升级后加载失败,怎么办
PyTorch 升版本之后加载旧模型失败,是很多人都会遇到的。如果你保存的是 state_dict,跨版本加载失败的情况其实不算多,但确实存在。常见原因有:
- 某些算子在不同版本里 module 的子模块命名发生了细微变化
- 旧版本 checkpoint 里存了不兼容的类引用(比如保存了完整模型对象)
- buffer 的数量发生变化(比如某些版本的自注意力层从无位置编码改成了带位置编码)
一个比较通用的策略:加载时先打印 state_dict.keys() 和你模型的 model.state_dict().keys(),做一次差集分析。哪些 key 在文件里有但是模型里没有(Unexpected),哪些 key 在模型里有但是文件里没有(Missing)。
python复制model_keys = set(model.state_dict().keys())
loaded_keys = set(state_dict.keys())
print('missing:', model_keys - loaded_keys)
print('unexpected:', loaded_keys - model_keys)
这个方法比直接看报错信息有用得多,因为 PyTorch 的报错虽然会列出 missing 和 unexpected 的 key,但经常因为列表太长被省略。通过差集,你一眼就能看出是哪一层的结构变了。
6.2 自定义模型升级时的手动键名映射
项目迭代过程中,模型的层经常要改。比如之前 self.conv1 是个 5x5 卷积,后来你发现 3x3 效果更好,把代码改成了 self.conv1 仍然是 3x3。这种同名不同结构的更改还好,PyTorch 加载时只关注维度匹配,会自动报 size mismatch 错误。
但如果你把 self.fc1 改名为 self.classifier,那旧权重文件里 fc1.weight 就找不到对应的层了。此时有两种处理办法:
办法一:加载前手动构造一个映射字典
python复制old_state = torch.load('old_model.pt', map_location='cpu')
rename_map = {'fc1.weight': 'classifier.weight', 'fc1.bias': 'classifier.bias'}
new_state = {}
for k, v in old_state.items():
new_k = rename_map.get(k, k)
new_state[new_k] = v
model.load_state_dict(new_state)
办法二:只加载能匹配的部分
用 strict=False 加载,然后对未匹配的层做随机初始化或 Kaiming 初始化。说白了,这就是 fine-tune 里"加载你能加载的,重新训练不能加载的"思路。
我个人偏向用映射字典,因为显式、可控,不会像 strict=False 那样容易掩盖真正的问题。
6.3 二进制权重格式转换的那些事
热词里有人搜"pytorch bin转换为pt",这里的 bin 一般是指某些预训练模型提供的权重格式,比如 Transformers 库的 pytorch_model.bin。它本质上就是 OrderedDict 的 pickle 序列化文件,只是后缀名不一样。
处理方式特别简单,其实 bin 和 pt 的内容是一回事,直接 torch.load 就能读出来。如果真要做"格式转换",通常指的是把权重加载出来,重新按你自己的 key 命名规则打包:
python复制state_dict = torch.load('pytorch_model.bin', map_location='cpu')
# 例如:huggingface 模型的 key 是 bert.embeddings.word_embeddings.weight
# 你希望变成 word_embeddings.weight,则自行构造映射
my_state_dict = {k.replace('bert.', ''): v for k, v in state_dict.items()}
torch.save(my_state_dict, 'converted_model.pt')
这里要提醒一下:除非你非常清楚 HuggingFace 模型的 state_dict 结构和你自定义模型的差异,否则不要盲目做全局 key 替换。最好的方式还是直接用对应库的 from_pretrained 接口加载,再提取 state_dict。
7. 实际工程中的几个高频坑与排查链路
7.1 加载到一半内存爆掉:图像模型 OOM
我在给一个图像分类项目做推理服务时遇到过:模型用 ResNet 在 GPU 上训练,checkpoint 每个 epoch 都存;到了部署环境,推理机只有 CPU,内存 16GB。加载的时候 map_location='cpu' 是写了,但程序直接 killed。
排查链路:
- 先用
nvidia-smi确认部署机没有 GPU。 free -h看内存确实紧张,系统可用只有 1.2GB。- checkpoint 文件本身 400MB,state_dict 里的 tensor 展开后占用远大于磁盘大小。
- 进一步查发现里面有一个缓存 list:训练时验证集每批的输出都 append 进了 checkpoint,导致 state_dict 没什么问题,是文件内容太大。
最终方案:重写保存逻辑,只保存模型参数、优化器状态和必要的字段,不再保存验证集临时数据。加载后内存占用降到 250MB。
7.2 Batchnorm size mismatch:输入尺寸改变导致加载失败
有一次我在做迁移学习,把 224x224 训练的模型换成 384x384 输入重新 fine-tune,结果加载时报错:
code复制size mismatch for encoder.blocks.0.attn.qkv.weight: copying a param with shape torch.Size([768, 768]) from checkpoint, the shape in current model is torch.Size([768, 768]).
看起来一模一样?那是因为我网络结构里某些层虽然一样,但输入序列长度变了,导致自适应池化层输出的维度变了,进而影响后续全连接层的输入维度。这种情况没有简单的映射办法,只能:
- 加载时跳过不匹配的层(用 strict=False)
- 在配置文件中记录模型的输入尺寸,尽量保持训练和推理的输入尺寸一致
踩过这次坑之后,我在项目规范里加了一条:所有模型的 checkpoint 保存时,必须同时保存一个 config.json,记录输入尺寸、归一化参数、类别数。加载的时候先检查 config 和当前模型是否一致,不一致就直接拒绝加载,而不是静默报错。
7.3 行业 Trick:把模型权重命名约定写进代码
如果你的工程化程度比较高,我给你推荐一个习惯:在模型类的 __init__ 中定义好权重文件的命名标准,并且专门暴露一个 save_weights 和 load_weights 方法。这样你的团队里任何一个成员拿到模型对象,不需要搜索 torch.save 的调用位置,就知道该怎么保存和加载。
python复制def save_weights(self, path: str):
checkpoint = {
'state_dict': self.state_dict(),
'config': self.config,
}
torch.save(checkpoint, path)
@classmethod
def from_pretrained(cls, path: str, device='cpu'):
checkpoint = torch.load(path, map_location='cpu')
config = checkpoint['config']
model = cls(config)
model.load_state_dict(checkpoint['state_dict'])
model.to(device)
model.eval()
return model
这个模式在 NLP 和 CV 的开源项目里都很常见,核心思想是把"模型状态的存取"封装成模型自身的职责,避免每一处调用都复制粘贴一段加载逻辑,等出问题的时候,排查链路也会清晰很多。
8. 我这些年总结出来的保存加载最佳实践清单
最后把核心要点整理成一份可以直接抄作业的清单,这些经验是我踩了无数次坑之后才沉淀下来的:
保存时:
- 首选保存
model.state_dict(),而不是完整模型对象。 - checkpoint 文件保存为字典对象,至少包含
epoch、model_state_dict、optimizer_state_dict、best_metric等字段。 - 保存前把模型显式切换到 CPU,避免意外把 GPU tensor 写进文件。
- 用
best_model.pt+last_model.pt滚动保存的策略,避免覆盖唯一快照。 - 保存一份
config.json,记录模型结构参数、输入尺寸、归一化参数和数据集信息。
加载时:
- 永远显式传
map_location。目标设备不确定时,统一先'cpu',再手动.to(device)。 - 加载后立刻
model.eval(),别忘了 BN 和 Dropout 的行为差异。 - 如果要做断点续训,优化器、scheduler 都要一起恢复,并且从保存的 epoch 继续。
- 多卡模型加载单卡环境,处理
module.前缀问题,尽量在保存侧就剥干净。 - Distribute 训练时,先构造原生模型加载 state_dict,再包 DDP。
这些经验放到不同场景下可能会有微调,但核心思路是稳定的。模型保存和加载说起来只是 PyTorch 里两个 API 调用,但在实际项目中,能不能把这些边界情况考虑清楚,直接决定了你的训练和部署流程稳不稳定。
我自己现在不管做什么模型,上来第一件事就是先把 save/load 的工具函数写好,再开始写训练主循环。因为我知道,只要这套机制健壮,训练中任何一次中断都不可怕,随时可以从最近的 checkpoint 继续跑。反过来,模型训得再好,存不下来、读不回去,那一切都等于零。
