【问题标题】:About saving state_dict/checkpoint in a function(PyTorch)关于在函数中保存 state_dict/checkpoint (PyTorch)
【发布时间】:2021-10-01 09:46:35
【问题描述】:

我正在尝试实现以下函数来保存 model_state 检查点:

def train_epoch(self):
for epoch in tqdm.trange(self.epoch, self.max_epoch, desc='Train Epoch', ncols=100):
    self.epoch = epoch      # increments the epoch of Trainer
    checkpoint = {} # fixme: here checkpoint!!!
    # model_save_criteria = self.model_save_criteria
    self.train()
    if epoch % 1 == 0:
        self.validate(checkpoint) 
    checkpoint_latest = {
        'epoch': self.epoch,
        'arch': self.model.__class__.__name__,
        'model_state_dict': self.model.state_dict(),
        'optimizer_state_dict': self.optim.state_dict(),
        'model_save_criteria': self.model_save_criteria
    }
    checkpoint['checkpoint_latest'] = checkpoint_latest
    torch.save(checkpoint, self.model_pth)

以前我只是通过运行一个 for 循环来做同样的事情:

train_states = {}
for epoch in range(max_epochs):
    running_loss = 0
    time_batch_start = time.time()
    model.train()
    for bIdx, sample in enumerate(train_loader):
        ...
        train...
        validation...
        train_states_latest = {
          'epoch': epoch + 1,
          'model_state_dict': model.state_dict(),
          'optimizer_state_dict': optimizer.state_dict(),
          'model_save_criteria': chosen_criteria}
        train_states['train_states_latest'] = train_states_latest
        torch.save(train_states, FILEPATH_MODEL_SAVE)

有没有办法启动checkpoint={} 并在每个循环中更新它?或者 checkpoint={} 在每个时期都很好,因为模型本身持有 state_dict()。只是我每次都覆盖检查点。

【问题讨论】:

    标签: pytorch state-dict


    【解决方案1】:

    您可以通过简单地更改 FILEPATH_MODEL_SAVE 路径并让该路径包含有关纪元或迭代次数的信息来避免覆盖检查点。例如(获取您的原始代码),

    train_states = {}
    for epoch in range(max_epochs):
        running_loss = 0
        time_batch_start = time.time()
        model.train()
        for bIdx, sample in enumerate(train_loader):
            ...
            train...
            validation...
            train_states_latest = {
              'epoch': epoch + 1,
              'model_state_dict': model.state_dict(),
              'optimizer_state_dict': optimizer.state_dict(),
              'model_save_criteria': chosen_criteria}
            train_states['train_states_latest'] = train_states_latest
            
    
            # This is the code you can add
            FILEPATH_MODEL_SAVE = "Epoch{}batch{}model_weights.pth".format(epoch, bIdx)
            torch.save(train_states, FILEPATH_MODEL_SAVE)
    
    
    

    使用 torch.save 上方的这段新代码,您可以避免覆盖检查点。

    萨塔克

    【讨论】:

    • 但这将保存多个模型,确切地说,如果我训练 200 个 epoch 并且每个 epoch 有 60 个批次,那么将有 200x60 个模型。我共享的 for 循环代码已经可以在一个路径中保存两个状态(最佳和最新)。我只是想在我展示的函数中实现它。
    • 我不太了解您想要新代码做什么。所以我知道你已经得到它来保存最新和最好的重量。你的目标是什么新功能。
    • @sarthakjain 很抱歉没有澄清。我发布的for 循环已经完美运行。我想把它作为一个函数来实现。我编写/启动了一些代码。 checkpoint={} 是否正确让我感到困惑。
    猜你喜欢
    • 2020-10-17
    • 1970-01-01
    • 2021-08-29
    • 2019-12-15
    • 2020-11-10
    • 2020-05-13
    • 1970-01-01
    • 2019-03-27
    • 2022-10-24
    相关资源
    最近更新 更多