【问题标题】:exceptions when loading the checkpoint of a PyTorch NN model加载 PyTorch NN 模型的检查点时出现异常
【发布时间】:2020-04-14 09:52:24
【问题描述】:

调用以下单元格中定义的函数时,抛出异常'TypeError: forward() 接受 2 个位置参数,但给出了 9 个' 这个document 提供了更多细节

def load_checkpoint(chkptJP):
checkpoint = torch.load(chkptJP)
model2 = model1(checkpoint['input_size'],
              checkpoint['output_size'],
              checkpoint['fc1'],
              checkpoint['fc2'],
              checkpoint['optimizer_state_dict'],
              checkpoint['epoch'],
              checkpoint['class_to_idx'],
              checkpoint['learning_rate'])
model2.load_state_dict(checkpoint['state_dict'])
return model2

写出检查点的代码如下:

checkpoint ={'input_size':512,
         'output_size':102,
         'fc1':256,
         'fc2':102,
         'state_dict': model.state_dict(),
         'optimizer_state_dict': optimizer.state_dict(),
         'epoch': epoch+1,
         'class_to_idx': model.class_to_idx,
         'learning_rate': 0.003}
torch.save(checkpoint,chkptJP)

【问题讨论】:

  • 您要解决的问题是什么?
  • 从共享的文档中,我得出的结论是,您正在使用自定义数据集对预训练模型进行迁移学习,并调整几层,然后对其进行检查点。是吗?
  • 嗨,你能显示转发功能代码吗?
  • 是的,它确实是从 RESNET18 开始的迁移学习,我通过替换和训练分类器来自定义:github.com/joepareti54/image-classification/blob/master/…

标签: python deep-learning pytorch checkpoint


【解决方案1】:

您的错误表明model1 是一个已经实例化的网络,而它应该是一个类。有关全面信息,请参阅official documentation about saving(如有疑问,请始终参考)。我将在整个答案中链接到它,因此请务必查看并了解发生了什么。

保存常规检查点

您的代码保存了general checkpoint。您可以通过这种方式保存任何字典以及您想要的任何信息(基本上是Python's pickle,您也可以类似地对其进行调整)。您的信息很多,其中一些与模型本身无关。

加载常规检查点

正如您所做的那样,您可以通过torch.load 加载所有这些数据。由于您保存了state_dict(权重),而不是整个Model(代码的外观),因此您必须使用随机权重创建模型并在之后加载它们。

这段代码应该没问题:

def load_checkpoint(chkptJP):
    checkpoint = torch.load(chkptJP)
    model = ModelClass(
        checkpoint["input_size"],
        checkpoint["output_size"],
        checkpoint["fc1"],
        checkpoint["fc2"],
        checkpoint["optimizer_state_dict"],
        checkpoint["epoch"],
        checkpoint["class_to_idx"],
        checkpoint["learning_rate"],
    )
    model.load_state_dict(checkpoint["state_dict"])
    return model

注意 ModelClass 必须是类,而不是您在此处所做的对象。如果model1 是一个对象,运行model1(arg1, ..., arg9) 将调用它的__call__ 方法,如果model1torch.nn.Module 的一个实例,则该方法又是一个包装的forward 方法。 ModelClass 在您的代码中应该是这样的(并且可能在某处定义):

import torch


class ModelClass(torch.nn.Module):
    def __init__(
        self,
        input_size,
        output_size,
        fc1,
        fc2,
        optimizer_state_dict,
        epoch,
        class_to_idx,
        learning_rate,
    ):
        # Your initialization code here
        ...

    def forward(tensor):
        # Your forward pass here
        ...

如果您在任何地方都没有 ModelClass,则必须单独保存整个模型(例如,torch.save(model) 而不是 torch.save(model.state_dict())) 并将其作为一个整体加载(torch.load(PATH) 而不是 chkp=torch.load(PATH),然后是 model.load_state_dict在实例上调用)

【讨论】:

  • 非常感谢;就我而言,我使用的是 torchvision 的 resnet18。我在任何地方都看不到 ModelClass
  • @josephpareti 在这种情况下是torchvision.models.resnet18
  • 不确定如何实施,抱歉。这里有一个 ResNet 类定义github.com/pytorch/vision/blob/master/torchvision/models/… 是否应该包含在我的笔记本中,鉴于我已经在使用 model1 = models.resnet18(pretrained=True),它是如何工作的
猜你喜欢
  • 2021-01-15
  • 2019-07-07
  • 2021-11-08
  • 2022-08-13
  • 2021-11-29
  • 1970-01-01
  • 2021-10-05
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多