【问题标题】:Loading model from checkpoint is not working从检查点加载模型不起作用
【发布时间】:2020-11-24 08:12:19
【问题描述】:

我训练了一个我从 this repository 修改的香草 vae。当我尝试使用经过训练的模型时,我无法使用 load_from_checkpoint 加载权重。我的检查点对象和lightningModule 对象之间似乎不匹配。

我已经使用pytorch-lightning LightningModule 设置了一个实验 (VAEXperiment)。我尝试将权重加载到网络中:

#building a new model
model = VanillaVAE(**config['model_params'])
model.build_layers()

#loading the weights
experiment = VAEXperiment(model, config['exp_params'])
experiment.load_from_checkpoint(path_to_checkpoint, config['exp_params'])

我也试过了:

checkpoint = torch.load(path_to_checkpoint, map_location=lambda storage, loc: storage)
model.load_state_dict(checkpoint['state_dict'])

但我得到一个错误 Unexpected key(s) in state_dict: "model.encoder.0.0.weight", "model.encoder.0.0.bias"...

我也关注了这个问题 https://github.com/PyTorchLightning/pytorch-lightning/issues/924 https://github.com/PyTorchLightning/pytorch-lightning/issues/2798

为什么我会收到此错误?是因为我的模型中的编码器和解码器模块吗?根据 git 上的问题日志,似乎错误已解决。我做错了什么?

【问题讨论】:

  • 后一种情况的问题是VanillaVAE.model.encoder不存在。但是VanillaVAE.encoder 可以。你试过experiment.load_state_dict(checkpoint['state_dict'])吗?
  • 谢谢罗曼。这是正确的答案。不敢相信我没弄明白。
  • 为什么第一种方法 (experiment.load_from_checkpoint) 失败了?实际上,它在我的代码中也失败了。但是,第二种方法是有效的。

标签: pytorch pytorch-lightning


【解决方案1】:

发布来自 cmets 的答案:

experiment.load_state_dict(checkpoint['state_dict'])

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2021-10-21
    • 2019-11-18
    • 2019-06-13
    • 2021-04-21
    • 1970-01-01
    • 2021-01-15
    • 2023-01-01
    • 1970-01-01
    相关资源
    最近更新 更多