Commit c49c292c authored by Wei Ho's avatar Wei Ho Committed by Facebook Github Bot
Browse files

Add CheckpointManager to keep avg checkpoint weights in memory to reduce disk...

Add CheckpointManager to keep avg checkpoint weights in memory to reduce disk read when averaging + various checkpoint refactoring

Summary: Pull Request resolved: https://github.com/pytorch/translate/pull/315

Reviewed By: akinh

Differential Revision: D13510446

fbshipit-source-id: 22a6594af9253130a93e638285a47183a974e0de
parent 829bd8ce
Loading
Loading
Loading
Loading
+1 −1
Original line number Diff line number Diff line
@@ -119,7 +119,7 @@ class Trainer(object):
        if distributed_utils.is_master(self.args):  # only save one checkpoint
            extra_state['train_meters'] = self.meters
            utils.save_state(
                filename, self.args, self.get_model(), self.criterion, self.optimizer,
                filename, self.args, self.get_model().state_dict(), self.criterion, self.optimizer,
                self.lr_scheduler, self._num_updates, self._optim_history, extra_state,
            )

+2 −2
Original line number Diff line number Diff line
@@ -39,7 +39,7 @@ def convert_state_dict_type(state_dict, ttype=torch.FloatTensor):
        return state_dict


def save_state(filename, args, model, criterion, optimizer, lr_scheduler,
def save_state(filename, args, model_state_dict, criterion, optimizer, lr_scheduler,
               num_updates, optim_history=None, extra_state=None):
    if optim_history is None:
        optim_history = []
@@ -47,7 +47,7 @@ def save_state(filename, args, model, criterion, optimizer, lr_scheduler,
        extra_state = {}
    state_dict = {
        'args': args,
        'model': model.state_dict() if model else {},
        'model': model_state_dict if model_state_dict else {},
        'optimizer_history': optim_history + [
            {
                'criterion_name': criterion.__class__.__name__,