From e57f4613231f29625448ebcfb65bfaceb4190f2e Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 15 Jul 2019 13:17:38 -0400 Subject: [PATCH] made checkpoint callback optional --- pytorch_lightning/models/trainer.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 81d27554..052a7621 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -31,7 +31,8 @@ class Trainer(TrainerIO): def __init__(self, experiment, - checkpoint_callback, early_stop_callback, + early_stop_callback, + checkpoint_callback=None, gradient_clip=0, cluster=None, process_position=0, @@ -68,7 +69,7 @@ class Trainer(TrainerIO): self.process_position = process_position self.current_gpu_name = current_gpu_name self.checkpoint_callback = checkpoint_callback - self.checkpoint_callback.save_function = self.save_checkpoint + self.checkpoint_callback.save_function = self.save_checkpoint if self.checkpoint_callback is not None else None self.early_stop = early_stop_callback self.model = None self.max_nb_epochs = max_nb_epochs @@ -720,5 +721,6 @@ class Trainer(TrainerIO): # model checkpointing if self.proc_rank == 0: - print('save callback...') - self.checkpoint_callback.on_epoch_end(epoch=self.current_epoch, logs=self.__tng_tqdm_dic) + if self.checkpoint_callback: + print('save callback...') + self.checkpoint_callback.on_epoch_end(epoch=self.current_epoch, logs=self.__tng_tqdm_dic)