From 8fd7a6001bdbe06159b1755ad8df77c450e6a789 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 24 Jul 2019 11:14:19 -0400 Subject: [PATCH] added safeguards for callbacks in loading saving --- pytorch_lightning/root_module/model_saving.py | 21 +++++++++++++------ 1 file changed, 15 insertions(+), 6 deletions(-) diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index a8bf366c..ab2ceb07 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -51,14 +51,19 @@ class TrainerIO(object): torch.save(checkpoint, filepath) def dump_checkpoint(self): + checkpoint = { 'epoch': self.current_epoch, - 'checkpoint_callback_best': self.checkpoint_callback.best, - 'early_stop_callback_wait': self.early_stop_callback.wait, - 'early_stop_callback_patience': self.early_stop_callback.patience, 'global_step': self.global_step } + if self.checkpoint_callback is not None: + checkpoint['checkpoint_callback_best'] = self.checkpoint_callback_best.best + + if self.early_stop_callback is not None: + checkpoint['early_stop_callback_wait'] = self.early_stop_callback.wait + checkpoint['early_stop_callback_patience'] = self.early_stop_callback.patience + optimizer_states = [] for i, optimizer in enumerate(self.optimizers): optimizer_states.append(optimizer.state_dict()) @@ -104,9 +109,13 @@ class TrainerIO(object): :param checkpoint: :return: """ - self.checkpoint_callback.best = checkpoint['checkpoint_callback_best'] - self.early_stop_callback.wait = checkpoint['early_stop_callback_wait'] - self.early_stop_callback.patience = checkpoint['early_stop_callback_patience'] + if self.checkpoint_callback is not None: + self.checkpoint_callback.best = checkpoint['checkpoint_callback_best'] + + if self.early_stop_callback is not None: + self.early_stop_callback.wait = checkpoint['early_stop_callback_wait'] + self.early_stop_callback.patience = checkpoint['early_stop_callback_patience'] + self.global_step = checkpoint['global_step'] self.current_epoch = checkpoint['epoch']