You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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报错。

修复步骤

  1. 重构Sequential定义,每次使用激活层时都实例化新的对象,避免复用同一个模块实例
  2. 拆分模型/优化器的实例化和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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 04:48:05