From 423bc5c6c9217dff44244731717b13a3d8252576 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 24 Jul 2019 18:01:33 -0400 Subject: [PATCH] testing hpc save load --- pytorch_lightning/root_module/model_saving.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index 1e47e3e7..493aa05b 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -144,7 +144,8 @@ class TrainerIO(object): filepath = '{}/hpc_ckpt_{}.ckpt'.format(folderpath, ckpt_number) # give model a chance to do something on hpc_save - self.model.on_hpc_save() + model = self.model.module if type(self.model) is LightningDataParallel else self.model + model.on_hpc_save() # request what to save from the model checkpoint_dict = self.dump_checkpoint() @@ -168,7 +169,7 @@ class TrainerIO(object): model.load_model_specific(checkpoint) # call model hook - self.model.on_hpc_load() + model.on_hpc_load() def max_ckpt_in_folder(self, path): files = os.listdir(path)