保存模型应只用torch.save(model.state_dict(), path),加载时先实例化模型再调用load_state_dict(torch.load(path)),注意module.前缀、device一致性及类定义匹配。

保存state_dict时必须用torch.save(),不能直接pickle模型对象
直接pickle.dump(model, f)或torch.save(model, path)会把整个模型类、构造参数、甚至Python闭包一起序列化,导致后续加载时严重依赖原始代码结构和环境。一旦类定义变了、路径改了、PyTorch版本升级,torch.load()就大概率报AttributeError: 'dict' object has no attribute 'forward'或ModuleNotFoundError。
正确做法是只保存可迁移的权重数据——即state_dict()返回的OrderedDict:
torch.save(model.state_dict(), "model_weights.pth")
这个文件体积小、不绑定类实现、跨机器/环境稳定。
加载前必须先实例化模型,再用load_state_dict()注入权重
torch.load()返回的是纯字典,不是模型。跳过模型实例化直接加载会报TypeError: load_state_dict() missing 1 required positional argument: 'state_dict'。
必须按原结构重建模型(哪怕只是空壳),再调用load_state_dict():
model = MyNet() # 必须和保存时完全一致的类定义
model.load_state_dict(torch.load("model_weights.pth"))
常见疏漏点:
- 模型类在加载脚本里没import,或import路径不同 →
NameError - 模型构造函数参数和训练时不一致(比如
num_classes=10写成5)→RuntimeError: size mismatch - 忘记调用
model.eval(),导致BatchNorm/ Dropout行为异常 → 验证指标波动大
多GPU训练后保存的state_dict含module.前缀,单卡加载需处理
用nn.DataParallel或DistributedDataParallel训练时,model.state_dict()的key默认带module.前缀(如"module.conv1.weight")。但单卡模型定义没有这层包装,直接load_state_dict()会因key不匹配失败,报Unexpected key(s) in state_dict。
两种解法:
- 加载时剔除前缀:
state_dict = torch.load("model.pth")<br>state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()}<br>model.load_state_dict(state_dict) - 训练时统一用
model.module.state_dict()保存(仅限DataParallel):torch.save(model.module.state_dict(), "model.pth")
注意:DistributedDataParallel不推荐用model.module.state_dict(),因其封装更复杂;优先选第一种清洗方式。
保存/加载时务必检查device一致性,避免CPU/GPU错配
在GPU上训练并保存的state_dict,如果直接在CPU上加载,load_state_dict()不会自动转设备,会报RuntimeError: Expected all tensors to be on the same device(尤其当模型部分参数在GPU、部分在CPU时)。
安全做法是显式指定map_location:
# 加载到CPU<br>model.load_state_dict(torch.load("model.pth", map_location="cpu"))<br><br># 加载到指定GPU<br>model.load_state_dict(torch.load("model.pth", map_location="cuda:1"))
反过来,如果模型在CPU上定义但想加载GPU权重,也得用map_location,否则torch.load()默认把tensor放回原设备(可能已不存在)。
最容易被忽略的是:保存时没留意model.to(device)是否生效,导致state_dict()里混着CPU和GPU tensor —— 这种情况torch.save()虽不报错,但加载时必然崩。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











