1. PyTorch模型保存的两种姿势:从入门到避坑

刚训练好的PyTorch模型就像刚出炉的蛋糕,保存不当就会变质。新手最常问的问题是:"为什么我的模型换个电脑就打不开了?"这通常是因为没搞懂两种保存方式的本质区别。

第一种是完整模型保存法,相当于把蛋糕连带烤盘一起打包:

torch.save(model, 'model.pth')

这种方式简单粗暴,但隐患很大。我去年就踩过坑——把实验室服务器上训练的模型发给同事后,他那边一直报"ImportError: No module named 'custom_layers'"的错误。后来发现是因为保存时连带记录了Python路径信息,就像快递单上写着"必须送货到XX小区3单元"。

第二种是参数字典保存法,官方推荐的做法:

torch.save(model.state_dict(), 'params.pth')

这就像只保存蛋糕的配方和食材比例。我在实际项目中发现,用这种方法保存的VGG16模型文件大小只有528KB,而完整保存的要535KB。虽然只差7KB,但当你要部署到移动端时,这点差异可能决定APP能否通过应用商店审核。

注意:使用state_dict()保存时,模型结构定义代码需要单独保存。我习惯把模型类代码和参数文件放在同一目录,就像把食谱和食材打包在一起。

2. 模型加载的三大雷区与排雷指南

加载模型时最容易翻车的三个地方,我用血泪教训总结成了避坑清单:

2.1 设备不匹配的经典错误

当你用GPU训练的模型要在CPU上加载时,不加map_location参数就会报错。正确的打开方式:

device = 'cuda' if torch.cuda.is_available() else 'cpu'
model.load_state_dict(torch.load('params.pth', map_location=device))

上周帮学弟debug时发现,他在Colab上训练后直接torch.load(),回到自己笔记本上就崩溃了。加上map_location后就像装了万能适配器,自动处理设备差异。

2.2 结构未定义的常见陷阱

只加载参数时,必须提前定义好模型结构。这就像拿到乐高零件包,但没有说明书就拼不出原造型。我建议采用这种安全模式:

# 先重建模型结构(建议单独保存在model.py)
from model import MyModel  
model = MyModel()
# 再加载参数
model.load_state_dict(torch.load('params.pth'))

2.3 版本兼容性的隐藏炸弹

PyTorch不同版本间可能存在兼容问题。有次我用1.8训练的模型,在1.6环境加载时报了奇怪的shape不匹配错误。解决方法是在保存时加上_use_new_zipfile_serialization:

torch.save(model.state_dict(), 'params.pth', 
          _use_new_zipfile_serialization=False)

3. 文件后缀的玄学真相

.pth、.pt、.ckpt这些后缀到底有什么区别?实测发现它们就像不同颜色的U盘——存储内容完全一样。我用ResNet18做了组对照实验:

后缀类型 文件大小 可加载性
.pth 44.7MB
.pt 44.7MB
.ckpt 44.7MB
.bin 44.7MB

虽然官方示例常用.pth,但我在开源项目里更常见到.pt。有个趣事:某次提交代码时用了.model后缀,review时被组长吐槽"你这扩展名太有创意了"。

4. 模型结构查看的六种武器

想知道模型里面长什么样?这几个方法比X光还好用:

4.1 直接打印法

print(model)

输出像解剖图一样层层展开,但遇到复杂模型时可能刷屏。有次打印Transformer模型,控制台直接滚了300多行...

4.2 逐层扫描术

for name, layer in model.named_children():
    print(f"{name}: {layer}")

这就像用显微镜观察细胞结构。调试时发现某个卷积层异常时,可以精准定位到"features.12.conv3"这样的具体位置。

4.3 参数统计法

total_params = sum(p.numel() for p in model.parameters())
print(f"总参数量:{total_params:,}")

当老板问"这模型有多大"时,这个数字比说"大概几十兆"专业多了。实测VGG16有138,357,544个参数,1亿3千多万!

4.4 可视化工具链

from torchsummary import summary
summary(model, input_size=(3, 224, 224))

输出结构化表格,包含每层的输出维度。记得第一次看到这个输出时,我突然理解了为什么输入图片要resize到224×224。

4.5 参数遍历技巧

for name, param in model.named_parameters():
    print(f"{name}: {param.shape}")

检查参数形状时特别有用。曾经发现某层的weight应该是[64,3,3,3]但实际是[64,1,3,3],这才知道前面层定义错了。

4.6 张量流追踪

x = torch.randn(1, 3, 224, 224)
for layer in model.children():
    x = layer(x)
    print(x.shape)

像给模型做胃肠镜,看到数据在每个层的变形过程。有个项目里就这样发现了某个MaxPool层输出意外变成了[1,512,6,6]。

5. 工程化实践中的生存法则

在真实项目中,这些经验能让你少加班:

  1. 版本控制黄金组合:把model.py和params.pth一起git管理,就像保存源代码和二进制
  2. 模型验证必做步骤:加载后先用测试数据跑forward,我习惯加个assert检查输出shape
  3. 云存储注意事项:当模型>100MB时,建议用torch.save的压缩选项:
    torch.save(model.state_dict(), 'model.zip', 
              _use_new_zipfile_serialization=True)
    
  4. 安全加载规范:从不可信来源加载模型时,一定要用pickle的安全加载:
    model = torch.load('model.pth', pickle_module=dill)
    

有次团队协作时,A同事的模型在B电脑上始终报错,最后发现是B的Python环境缺少某个科学计算库。现在我们的做法是用Docker容器打包整个环境,像罐头一样密封交付。

模型部署到生产环境时,还要考虑转换为TorchScript格式。虽然本教程不涉及这个进阶话题,但记住这个转换命令能救急:

script_model = torch.jit.script(model)
script_model.save('model.pt')

最后说个真实案例:某次模型验证准确率突然从90%跌到随机水平,查了三天发现是有人误用了torch.save(model)而不是model.state_dict(),导致加载时结构被意外修改。所以再强调一次——保存参数字典是最佳实践

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐