【问题标题】:How to use checkpoint model file in pytorch to test the CIFAR-10 dataset?如何在 pytorch 中使用 checkpoint 模型文件来测试 CIFAR-10 数据集?
【发布时间】:2019-03-27 00:43:53
【问题描述】:
model = SqueezeNext()
model = model.to(device)

def load_checkpoint(model, optimizer, losslogger, filename='SqNxt_23_1x_Cifar.ckpt'):
# Note: Input model & optimizer should be pre-defined.  This routine only updates their states.
start_epoch = 0
if os.path.isfile(filename):
    print("=> loading checkpoint '{}'".format(filename))
    checkpoint = torch.load(filename)
    start_epoch = checkpoint['epoch']
    model.load_state_dict(checkpoint['state_dict'])
    optimizer.load_state_dict(checkpoint['optimizer'])
    losslogger = checkpoint['losslogger']
    print("=> loaded checkpoint '{}' (epoch {})"
              .format(filename, checkpoint['epoch']))
else:
    print("=> no checkpoint found at '{}'".format(filename))


return model, optimizer, start_epoch, losslogger

model, optimizer, start_epoch, losslogger = load_checkpoint(model, optimizer, losslogger)

TypeError: Traceback(最近一次调用最后一次) 在 () 41 test_loader = torch.utils.data.DataLoader(test_dataset,batch_size=80,num_workers=8,shuffle=False) 42 ---> 43 模型 = SqueezeNext() 44 模型 = 模型.to(设备) 45 def load_checkpoint(model, optimizer, losslogger, filename='SqNxt_23_1x_Cifar.ckpt'): TypeError: init() 缺失 3 所需的位置参数:“width_x”、“blocks”和“num_classes”

我认为我没有以正确的方式实现这一点!

【问题讨论】:

    标签: python python-3.x deep-learning pytorch torchvision


    【解决方案1】:

    您的错误不是来自您的检查点功能。如果我们查看回溯:

    > TypeError: Traceback (most recent call last)
    > <ipython-input-51-94c8be648862> in <module>()
    >      41 test_loader   = torch.utils.data.DataLoader(test_dataset, batch_size=80, num_workers=8, shuffle=False)
    >      42 
    > ---> 43 model = SqueezeNext()
    >      44 model = model.to(device)
    >      45 def load_checkpoint(model, optimizer, losslogger, filename='SqNxt_23_1x_Cifar.ckpt'): TypeError: __init__() missing 3
    > required positional arguments: 'width_x', 'blocks', and 'num_classes'
    

    我们被告知的行是第 43 行:

    > ---> 43 model = SqueezeNext()
    

    错误是:

    > required positional arguments: 'width_x', 'blocks', and 'num_classes'
    

    我假设您使用的是 SqueezeNext 的 this implementation,但无论您使用哪种实现,您都没有传递初始化模型所需的所有参数。您需要将代码更改为:

    model = SqueezeNext(width_x=1.0, blocks=[6, 6, 8, 1], num_classes=10)
    

    如果您不使用该实现,则需要找到 SqueezeNext 模型的源代码,并查看 __init__ 函数需要哪些参数。你可以试试这个:

    import inspect
    
    inspect.signature(SqueezeNext.__init__)
    

    应该给你签名。

    【讨论】:

    • @RMPR 感谢您的编辑。但请注意,“initialise”不是拼写错误,而是正确的英式拼写。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2018-06-22
    • 2021-08-29
    • 2021-05-01
    • 2016-06-19
    • 2020-04-25
    • 2015-11-13
    • 2022-12-04
    相关资源
    最近更新 更多