加载PyTorch .pth模型用于推理时state_dict键不匹配如何解决?
错误原因
- 核心问题是你当前实例化的
GatherModel结构和保存cigin.tar权重时使用的模型结构完全不匹配:你存的权重文件里包含的是solute_pass、solvent_pass、独立溶质/溶剂LSTM层、四层全连接层的参数,而你现在写的GatherModel中定义的是lin0、set2set、message_layer、conv这类层,两边的层命名、结构设计完全对应不上,才会同时出现大量缺失键和意外键的报错。 - 你添加
strict=False参数后看到的不是报错,是PyTorch返回的不匹配键提示,这个参数本身只是跳过键不匹配的报错逻辑,并没有解决结构不匹配的问题,此时你的模型所有层都没有加载到正确的权重,推理结果完全不可用。
解决方法
- 优先使用权重对应的原始模型定义:你用的
cigin.tar是CIGIN模型的预训练权重,直接找到该权重开源对应的原始模型类代码实例化,再加载权重即可消除不匹配问题,这是成本最低的解决方案。 - 如果你确实需要用当前自定义的
GatherModel结构,需要手动做权重键映射:先打印出加载的原始权重的所有键,核对每个权重参数对应你自定义模型里的哪一层的参数,构建键名映射表,把原始权重的键名替换成你模型里的键名后再加载,示例代码如下:
import torch from your_module import GatherModel # 实例化自定义模型 model = GatherModel() # 加载原始权重文件 raw_weights = torch.load('/content/CIGIN/weights/cigin.tar') # 自定义键映射表,需要你根据参数的形状、作用核对对应关系后补全所有项 key_map = { "solute_pass.U_0.weight": "lin0.weight", "solute_pass.U_0.bias": "lin0.bias", # 剩余所有键的映射关系都需要补全 } # 生成适配当前模型的state_dict adapted_state_dict = {} for old_key, param in raw_weights.items(): if old_key in key_map: adapted_state_dict[key_map[old_key]] = param # 加载适配后的权重 model.load_state_dict(adapted_state_dict, strict=True)
- 额外检查:确认你有没有加载错权重文件,排除拿了其他任务、其他结构训练出来的权重的情况。
内容的提问来源于stack exchange,提问作者harsh
相关产品推荐
相关产品推荐

