PyTorch使用load_state_dict加载训练模型报错的正确处理方法咨询
PyTorch模型state_dict加载报错解决方案
问题原因
- 链式调用
load_state_dict()赋值错误:load_state_dict()返回的是_IncompatibleKeys类型的键校验结果对象,不是模型/优化器本身,直接赋值会导致后续调用时报_IncompatibleKeys不可调用的错误。 - Sequential结构定义不规范:代码中提前实例化了全局的
act=torch.nn.ReLU()对象,并在Sequential中多次复用同一个实例,PyTorch不会重复注册同一个模块实例,导致state_dict的层序号和预期不符,触发unexpected key报错。
修复步骤
- 重构Sequential定义,每次使用激活层时都实例化新的对象,避免复用同一个模块实例
- 拆分模型/优化器的实例化和state_dict加载步骤,先创建空白实例,再调用
load_state_dict()加载参数
修正后可运行代码
import torch lin=torch.nn.Linear fnc=torch.nn.functional class Ann(torch.nn.Module): def __init__(self): super(Ann, self).__init__() self.conv1 = torch.nn.Conv2d(1, 10, kernel_size=5) self.conv2 = torch.nn.Conv2d(10, 20, kernel_size=4) self.drop = torch.nn.Dropout2d(p=0.5) self.fc1 = torch.nn.Linear(320,128) self.fc2 = torch.nn.Linear(128,10) def forward(self, x): x = self.conv1(x[:,None,:,:]) x = fnc.relu(fnc.max_pool2d(x,2)) x = self.drop(self.conv2(x)) x = fnc.relu(fnc.max_pool2d(x,2)) x = torch.flatten(x,1) x = fnc.relu(self.fc1(x)) x = fnc.dropout(self.fc2(x),training=self.training) return fnc.log_softmax(x,dim=1) x,y=torch.rand((5,28,28)),torch.randint(0,9,(5,)) f=fnc.nll_loss # 修正:不要复用同一个ReLU实例,每次实例化新的激活层 ann1 = torch.nn.Sequential( torch.nn.Flatten(start_dim=1), lin(784,256), torch.nn.ReLU(), lin(256,128), torch.nn.ReLU(), lin(128,10), torch.nn.LogSoftmax(dim=1) ) ann2=Ann() F1 = torch.optim.SGD(ann1.parameters(),lr=0.01,momentum=0.5) F2 = torch.optim.SGD(ann2.parameters(),lr=0.01,momentum=0.5) # 训练步骤保持不变 F1.zero_grad(); y_=ann1(x); loss=f(y_,y); loss.backward(); F1.step() print(x.dtype,y.dtype,x.shape,y.shape,y_.shape,loss) F2.zero_grad(); y_=ann2(x); loss=f(y_,y); loss.backward(); F2.step() print(x.dtype,y.dtype,x.shape,y.shape,y_.shape,loss) name='/home/leon/' # 保存逻辑不变 torch.save([ann1.state_dict(),F1.state_dict()], name+'annF1.pth') torch.save([ann2.state_dict(),F2.state_dict()], name+'annF2.pth') a1,d1=torch.load(name+'annF1.pth') a2,d2=torch.load(name+'annF2.pth') # 修正加载逻辑:先实例化,再加载参数 ann3 = type(ann1)(*ann1.args) ann3.load_state_dict(a1) F3 = torch.optim.SGD(ann3.parameters(),lr=0.01,momentum=0.5) F3.load_state_dict(d1) ann4 = type(ann2)() ann4.load_state_dict(a2) F4 = torch.optim.SGD(ann4.parameters(),lr=0.01,momentum=0.5) F4.load_state_dict(d2) # 测试加载后的模型是否正常运行 print(ann3(x)) print(ann4(x))
额外说明
原自定义Ann类的forward方法中log_softmax的dim参数默认设置为0,是按样本维度做归一化,不符合分类任务常规的按输出维度归一化的逻辑,会导致loss计算异常,上述代码中已调整为dim=1。
内容的提问来源于stack exchange,提问作者Leo
相关产品推荐
相关产品推荐

