【问题标题】:How to load a saved model defined in a function in PyTorch in Colab?如何在 Colab 的 PyTorch 中加载函数中定义的保存模型?
【发布时间】:2021-08-17 21:41:21
【问题描述】:

这是我的训练功能的示例代码(删除了不必要的部分):

我正在尝试将我的模型 data_gen 保存在 torch.save() 中,运行 train_dmc 函数后,我可以在目录中找到检查点文件。

def train_dmc(loader,loss):


 
  data_gen = DataGenerator().to(device)

  data_gen_optimizer = optim.Rprop(para_list, lr=lrate)


  savepath='/content/drive/MyDrive/'+loss+'checkpoint.t7'
  state = {
            'epoch': epoch,
            'model_state_dict': data_gen.state_dict(),
            'optimizer_state_dict': data_gen_optimizer.state_dict(),
            'data loss': data_loss,
            'latent_loss':latent_loss
            }
  torch.save(state,savepath)

我的问题是,如果 Google Colab 断开连接,如何加载检查点文件以继续训练。

我应该加载 data_gen 还是 train_dmc(),这是我第一次使用它,我真的很困惑,因为 data_gen 是在另一个函数中定义的。希望有人能帮我解释一下

data_gen.load_state_dict(torch.load(PATH))
data_gen.eval()

#or

train_dmc.load_state_dict(torch.load(PATH))
train_dmc.eval()

【问题讨论】:

    标签: python pytorch torch checkpoint


    【解决方案1】:

    由于state 变量是字典,所以尝试将其保存为:

    with open('/content/checkpoint.t7', 'wb') as handle:
        pickle.dump(state, handle, protocol=pickle.HIGHEST_PROTOCOL)
    

    将您的模型类初始化为data_gen = DataGenerator().to(device)

    并将检查点文件加载为:

    import pickle
    file = open('/content/checkpoint.t7', 'rb')
    loaded_state = pickle.load(file)
    

    然后您可以使用data_gen = loaded_state['model_state_dict'] 加载state_dict。这会将 state_dict 加载到模型类!

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2021-12-15
      • 2019-09-26
      • 1970-01-01
      • 2020-05-03
      • 2020-07-17
      • 2019-08-30
      • 1970-01-01
      相关资源
      最近更新 更多