PyTorch Sequential加载state_dict报错:存在意外键0.weight等
PyTorch Sequential模型加载报错解决方案
问题原因
你创建的空Sequential模型和原模型结构完全不符,加载时state_dict里的键(如0.weight)对应原模型的层参数,但空模型没有这些层,因此触发Unexpected key错误。
解决步骤
1. 重建与原模型完全一致的结构
加载前必须先定义和训练时一模一样的Sequential模型结构,不能直接用空的Sequential()。
比如原训练时的模型定义是:
import torch.nn as nn # 原训练用的模型结构 model = nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 30), nn.ReLU(), nn.Linear(30, 5) )
加载时也要先定义相同结构:
# 先重建相同结构的模型 model = nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 30), nn.ReLU(), nn.Linear(30, 5) )
2. 加载state_dict
完成结构重建后,再执行加载命令:
import torch model.load_state_dict(torch.load(PATH))
优化建议
为避免重复编写模型结构,可将模型构建逻辑封装成函数,在保存和加载的代码中复用:
def build_sequential_model(): return nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 30), nn.ReLU(), nn.Linear(30, 5) ) # 保存模型时 model = build_sequential_model() # ...训练逻辑... torch.save(model.state_dict(), PATH) # 加载模型时 model = build_sequential_model() model.load_state_dict(torch.load(PATH))
内容的提问来源于stack exchange,提问作者Davide Maran
相关产品推荐
相关产品推荐

