diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 587b6465..05c3fc1a 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -244,7 +244,9 @@ class Trainer(TrainerIO): # ----------------------------- def fit(self, model): + # give model convenience properties model.trainer = self + model.experiment = self.experiment # transfer data loaders from model self.__get_dataloaders(model) diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index b3b178e0..3d150752 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -107,6 +107,9 @@ class TrainerIO(object): # save exp to make sure we get all the metrics experiment.save() + # close experiment to avoid issues + experiment.close() + ckpt_number = self.max_ckpt_in_folder(folderpath) + 1 if not os.path.exists(folderpath): diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 537ec277..2349e7e4 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -26,6 +26,7 @@ class LightningModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks): self.gradient_clip = hparams.gradient_clip self.trainer = None self.from_lightning = True + self.experiment = None # track if gpu was requested for checkpointing self.on_gpu = False