From 36f0b5bbd0dfcfc42819a9c595087c2fc3773334 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Hendrik=20Schr=C3=B6ter?= Date: Fri, 4 Oct 2019 19:35:02 +0000 Subject: [PATCH] Use getter instead of python property for the dataloaders (#275) * Use getter instead of python property for the dataloaders * Fix lint * Update trainer.py --- pytorch_lightning/root_module/decorators.py | 11 +++- pytorch_lightning/trainer/trainer.py | 67 +++++++++------------ tests/test_models.py | 32 ++++++---- 3 files changed, 57 insertions(+), 53 deletions(-) diff --git a/pytorch_lightning/root_module/decorators.py b/pytorch_lightning/root_module/decorators.py index 75854362..65062f73 100644 --- a/pytorch_lightning/root_module/decorators.py +++ b/pytorch_lightning/root_module/decorators.py @@ -10,13 +10,18 @@ def data_loader(fn): attr_name = '_lazy_' + fn.__name__ - @property - def _data_loader(self): + def _get_data_loader(self): try: value = getattr(self, attr_name) except AttributeError: try: value = fn(self) # Lazy evaluation, done only once. + if ( + value is not None and + not isinstance(value, list) and + fn.__name__ in['test_dataloader', 'val_dataloader'] + ): + value = [value] except AttributeError as e: # Guard against AttributeError suppression. (Issue #142) traceback.print_exc() @@ -25,4 +30,4 @@ def data_loader(fn): setattr(self, attr_name, value) # Memoize evaluation. return value - return _data_loader + return _get_data_loader diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 03fffe92..179cc089 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -24,6 +24,7 @@ from pytorch_lightning.utilities.debugging import MisconfigurationException import pdb from pytorch_lightning.trainer import ignored_warnings + try: from apex import amp APEX_AVAILABLE = True @@ -141,9 +142,9 @@ class Trainer(TrainerIO): self.nb_val_batches = 0 self.nb_training_batches = 0 self.nb_test_batches = 0 - self.train_dataloader = None - self.test_dataloader = None - self.val_dataloader = None + self.get_train_dataloader = None + self.get_test_dataloaders = None + self.get_val_dataloaders = None # training state self.model = None @@ -450,19 +451,21 @@ class Trainer(TrainerIO): def __layout_bookeeping(self): # determine number of training batches - self.nb_training_batches = len(self.train_dataloader) + self.nb_training_batches = len(self.get_train_dataloader()) self.nb_training_batches = int(self.nb_training_batches * self.train_percent_check) # determine number of validation batches # val datasets could be none, 1 or 2+ - if self.val_dataloader is not None: - self.nb_val_batches = sum(len(dataloader) for dataloader in self.val_dataloader) + if self.get_val_dataloaders() is not None: + self.nb_val_batches = sum(len(dataloader) for dataloader in self.get_val_dataloaders()) self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check) self.nb_val_batches = max(1, self.nb_val_batches) # determine number of test batches - if self.test_dataloader is not None: - self.nb_test_batches = sum(len(dataloader) for dataloader in self.test_dataloader) + if self.get_test_dataloaders() is not None: + self.nb_test_batches = sum( + len(dataloader) for dataloader in self.get_test_dataloaders() + ) self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check) self.nb_test_batches = max(1, self.nb_test_batches) @@ -481,10 +484,10 @@ class Trainer(TrainerIO): # make dataloader_idx arg in validation_step optional args = [batch, batch_idx] - if test and len(self.test_dataloader) > 1: + if test and len(self.get_test_dataloaders()) > 1: args.append(dataloader_idx) - elif not test and len(self.val_dataloader) > 1: + elif not test and len(self.get_val_dataloaders()) > 1: args.append(dataloader_idx) # handle DP, DDP forward @@ -530,9 +533,9 @@ class Trainer(TrainerIO): outputs = [] # run training - for dataloader_idx, dl in enumerate(dataloaders): + for dataloader_idx, dataloader in enumerate(dataloaders): dl_outputs = [] - for batch_idx, batch in enumerate(dl): + for batch_idx, batch in enumerate(dataloader): if batch is None: # pragma: no cover continue @@ -582,21 +585,11 @@ class Trainer(TrainerIO): :param model: :return: """ + self.get_train_dataloader = model.train_dataloader + self.get_test_dataloaders = model.test_dataloader + self.get_val_dataloaders = model.val_dataloader - self.train_dataloader = model.train_dataloader - self.test_dataloader = model.test_dataloader - self.val_dataloader = model.val_dataloader - - # handle returning an actual dataloader instead of a list of loaders - have_test_loaders = self.test_dataloader is not None - if have_test_loaders and not isinstance(self.test_dataloader, list): - self.test_dataloader = [self.test_dataloader] - - have_val_loaders = self.val_dataloader is not None - if have_val_loaders and not isinstance(self.val_dataloader, list): - self.val_dataloader = [self.val_dataloader] - - if self.use_ddp and not isinstance(self.train_dataloader.sampler, DistributedSampler): + if self.use_ddp and not isinstance(self.get_train_dataloader().sampler, DistributedSampler): msg = """ You're using multiple gpus and multiple nodes without using a DistributedSampler to assign a subset of your data to each process. To silence this warning, pass a @@ -615,8 +608,8 @@ class Trainer(TrainerIO): """ warnings.warn(msg) - if self.use_ddp and self.val_dataloader is not None: - for dataloader in self.val_dataloader: + if self.use_ddp and self.get_val_dataloaders is not None: + for dataloader in self.get_val_dataloaders(): if not isinstance(dataloader.sampler, DistributedSampler): msg = """ Your val_dataloader(s) don't use DistributedSampler. @@ -638,8 +631,8 @@ class Trainer(TrainerIO): warnings.warn(msg) break - if self.use_ddp and self.test_dataloader is not None: - for dataloader in self.test_dataloader: + if self.use_ddp and self.get_test_dataloaders is not None: + for dataloader in self.get_test_dataloaders(): if not isinstance(dataloader.sampler, DistributedSampler): msg = """ Your test_dataloader(s) don't use DistributedSampler. @@ -954,12 +947,12 @@ class Trainer(TrainerIO): # run tiny validation (if validation defined) # to make sure program won't crash during val ref_model.on_sanity_check_start() - if self.val_dataloader is not None and self.nb_sanity_val_steps > 0: + if self.get_val_dataloaders() is not None and self.nb_sanity_val_steps > 0: # reset progress_bar limit for sanity check if self.show_progress_bar: self.progress_bar.reset(self.nb_sanity_val_steps) - self.evaluate(model, self.val_dataloader, self.nb_sanity_val_steps, self.testing) + self.evaluate(model, self.get_val_dataloaders(), self.nb_sanity_val_steps, self.testing) # --------------------------- # CORE TRAINING LOOP @@ -970,8 +963,8 @@ class Trainer(TrainerIO): # run all epochs for epoch_nb in range(self.current_epoch, self.max_nb_epochs): # set seed for distributed sampler (enables shuffling for each epoch) - if self.use_ddp and hasattr(self.train_dataloader.sampler, 'set_epoch'): - self.train_dataloader.sampler.set_epoch(epoch_nb) + if self.use_ddp and hasattr(self.get_train_dataloader().sampler, 'set_epoch'): + self.get_train_dataloader().sampler.set_epoch(epoch_nb) # get model model = self.__get_model() @@ -1016,7 +1009,7 @@ class Trainer(TrainerIO): model.on_epoch_start() # run epoch - for batch_nb, batch in enumerate(self.train_dataloader): + for batch_nb, batch in enumerate(self.get_train_dataloader()): self.batch_nb = batch_nb self.global_step += 1 @@ -1325,12 +1318,12 @@ class Trainer(TrainerIO): model.on_pre_performance_check() # select dataloaders - dataloaders = self.val_dataloader + dataloaders = self.get_val_dataloaders() max_batches = self.nb_val_batches # calculate max batches to use if test: - dataloaders = self.test_dataloader + dataloaders = self.get_test_dataloaders() max_batches = self.nb_test_batches # cap max batches to 1 when using fast_dev_run diff --git a/tests/test_models.py b/tests/test_models.py index 62a7fef3..cdb59f74 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -128,7 +128,8 @@ def test_dp_resume(): dp_model = new_trainer.model dp_model.eval() - _ = [run_prediction(dataloader, dp_model, dp=True) for dataloader in trainer.val_dataloader] + for dataloader in trainer.get_train_dataloader(): + run_prediction(dataloader, dp_model, dp=True) # new model model = LightningTestModel(hparams) @@ -186,7 +187,7 @@ def test_running_test_pretrained_model_ddp(): new_trainer = Trainer(**trainer_options) new_trainer.test(pretrained_model) - run_prediction(model.test_dataloader, pretrained_model) + [run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()] # test we have good test accuracy clear_save_dir() @@ -784,7 +785,8 @@ def test_cpu_restore_training(): # if model and state loaded correctly, predictions will be good even though we # haven't trained with the new loaded model trainer.model.eval() - _ = [run_prediction(dataloader, trainer.model) for dataloader in trainer.val_dataloader] + for dataloader in trainer.get_val_dataloaders(): + run_prediction(dataloader, trainer.model) model.on_sanity_check_start = assert_good_acc @@ -852,8 +854,9 @@ def test_cpu_slurm_save_load(): # predict with trained model before saving # make a prediction - for batch in model.test_dataloader: - break + for dataloader in model.test_dataloader(): + for batch in dataloader: + break x, y = batch x = x.view(x.size(0), -1) @@ -968,8 +971,9 @@ def test_model_saving_loading(): assert result == 1, 'amp + ddp model failed to complete' # make a prediction - for batch in model.test_dataloader: - break + for dataloader in model.test_dataloader(): + for batch in dataloader: + break x, y = batch x = x.view(x.size(0), -1) @@ -1060,7 +1064,7 @@ def test_amp_gpu_ddp_slurm_managed(): pretrained_model = load_model(logger.experiment, save_dir, True) # test model preds - run_prediction(model.test_dataloader, pretrained_model) + [run_prediction(dataloader, pretrained_model) for dataloader in trainer.get_test_dataloaders()] if trainer.use_ddp: # on hpc this would work fine... but need to hack it for the purpose of the test @@ -1287,10 +1291,11 @@ def test_multiple_val_dataloader(): assert result == 1 # verify there are 2 val loaders - assert len(trainer.val_dataloader) == 2, 'Multiple val_dataloaders not initiated properly' + assert len(trainer.get_val_dataloaders()) == 2, \ + 'Multiple val_dataloaders not initiated properly' # make sure predictions are good for each val set - [run_prediction(dataloader, trainer.model) for dataloader in trainer.val_dataloader] + [run_prediction(dataloader, trainer.model) for dataloader in trainer.get_val_dataloaders()] def test_multiple_test_dataloader(): @@ -1318,10 +1323,11 @@ def test_multiple_test_dataloader(): result = trainer.fit(model) # verify there are 2 val loaders - assert len(trainer.test_dataloader) == 2, 'Multiple test_dataloaders not initiated properly' + assert len(trainer.get_test_dataloaders()) == 2, \ + 'Multiple test_dataloaders not initiated properly' # make sure predictions are good for each test set - [run_prediction(dataloader, trainer.model) for dataloader in trainer.test_dataloader] + [run_prediction(dataloader, trainer.model) for dataloader in trainer.get_test_dataloaders()] # run the test method trainer.test() @@ -1356,7 +1362,7 @@ def run_gpu_model_test(trainer_options, model, hparams, on_gpu=True): pretrained_model = load_model(logger.experiment, save_dir, on_gpu) # test new model accuracy - run_prediction(model.test_dataloader, pretrained_model) + [run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()] if trainer.use_ddp: # on hpc this would work fine... but need to hack it for the purpose of the test