Commit 0daba38e authored by Myle Ott's avatar Myle Ott
Browse files

Save and restore wall time in checkpoints

parent dc40ac58
Loading
Loading
Loading
Loading
+6 −6
Changes for fairseq/meters.py: 6 added lines, 6 removed lines.
Original line number Diff line number Diff line
@@ -28,10 +28,11 @@ class AverageMeter(object):

class TimeMeter(object):
    """Computes the average occurrence of some event per second"""
    def __init__(self):
        self.reset()
    def __init__(self, init=0):
        self.reset(init)

    def reset(self):
    def reset(self, init=0):
        self.init = init
        self.start = time.time()
        self.n = 0

@@ -40,12 +41,11 @@ class TimeMeter(object):

    @property
    def avg(self):
        delta = time.time() - self.start
        return self.n / delta
        return self.n / self.elapsed_time

    @property
    def elapsed_time(self):
        return time.time() - self.start
        return self.init + (time.time() - self.start)


class StopwatchMeter(object):
+2 −0
Changes for train.py: 2 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -273,6 +273,7 @@ def save_checkpoint(trainer, args, epoch, val_loss=None):
    extra_state = {
        'epoch': epoch,
        'val_loss': val_loss,
        'wall_time': trainer.get_meter('wall').elapsed_time,
    }

    if not args.no_epoch_checkpoints:
@@ -302,6 +303,7 @@ def load_checkpoint(args, trainer, train_dataloader):
            for i in range(epoch):
                _ = next(train_dataloader)
            epoch += 1
            trainer.get_meter('wall').reset(init=extra_state.get('wall_time', 0))
    return epoch