如何从torch.save保存的模型中恢复state_dict并解决加载报错
解决PyTorch模型state_dict键不匹配问题
问题分析
你遇到的报错核心是直接保存的模型对象的state_dict键结构,和新定义模型的键不匹配:原模型state_dict的键是0.weight这类(对应Sequential层内的层序号),而新模型里Sequential被封装在model1属性下,因此需要的键是model1.0.weight这种带前缀的格式。
解决步骤
加载原模型并提取state_dict
先加载你之前保存的模型对象,获取它的原始state_dict:import torch import torch.nn as nn # 加载原模型文件 model1 = torch.load('path') original_state_dict = model1.state_dict()手动修改state_dict的键
遍历原state_dict的所有键,给每个键添加model1.前缀,生成适配新模型的state_dict:new_state_dict = {} for key, tensor in original_state_dict.items(): # 为原键添加"model1."前缀 new_key = f"model1.{key}" new_state_dict[new_key] = tensor定义匹配的模型结构并加载修改后的state_dict
根据你提供的模型维度(输入2、隐藏层5神经元、输出1),定义对应的模型结构,再加载修改后的state_dict:class Model(nn.Module): def __init__(self): super().__init__() self.model1 = nn.Sequential( nn.Linear(2, 5), # 输入层到5神经元隐藏层 nn.ReLU(), # 激活函数(若原模型用了其他激活,替换成对应类型即可) nn.Linear(5, 1) # 隐藏层到输出层 ) def forward(self, x): return self.model1(x) # 实例化模型并加载适配后的state_dict pretrained_model = Model() pretrained_model.load_state_dict(new_state_dict)验证模型可用性
可以用测试张量验证模型是否正常运行:test_input = torch.randn(1, 2) # 符合输入维度的测试数据 output = pretrained_model(test_input) print(output.shape) # 应输出torch.Size([1, 1]),匹配输出维度要求
额外注意
如果原模型的Sequential层使用了其他激活函数(比如Sigmoid、Tanh),需要将代码中的nn.ReLU()替换为对应的激活函数,否则模型结构不匹配仍会报错。
内容的提问来源于stack exchange,提问作者Acad
相关产品推荐
相关产品推荐

