From 6111edaf824e70c74d95faafa293a9382a577df3 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sat, 19 Oct 2019 00:39:30 +0200 Subject: [PATCH] Test fx (#390) * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx * changes to test fx --- pytorch_lightning/trainer/trainer.py | 51 ++++++++++++++++------------ 1 file changed, 30 insertions(+), 21 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 984dbdcd..0bce3e96 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -597,6 +597,9 @@ class Trainer(TrainerIOMixin): model.zero_grad() model.eval() + # copy properties for forward overrides + self.__copy_trainer_model_properties(model) + # disable gradients to save memory torch.set_grad_enabled(False) @@ -842,6 +845,24 @@ class Trainer(TrainerIOMixin): self.__run_pretrain_routine(model) + def __copy_trainer_model_properties(self, model): + if isinstance(model, LightningDataParallel): + ref_model = model.module + elif isinstance(model, LightningDistributedDataParallel): + ref_model = model.module + else: + ref_model = model + + for m in [model, ref_model]: + m.trainer = self + m.on_gpu = self.on_gpu + m.use_dp = self.use_dp + m.use_ddp2 = self.use_ddp2 + m.use_ddp = self.use_ddp + m.use_amp = self.use_amp + m.testing = self.testing + m.single_gpu = self.single_gpu + def ddp_train(self, gpu_nb, model): """ Entry point into a DP thread @@ -890,13 +911,7 @@ class Trainer(TrainerIOMixin): model.cuda(gpu_nb) # set model properties before going into wrapper - model.trainer = self - model.on_gpu = self.on_gpu - model.use_dp = self.use_dp - model.use_ddp2 = self.use_ddp2 - model.use_ddp = self.use_ddp - model.use_amp = self.use_amp - model.testing = self.testing + self.__copy_trainer_model_properties(model) # override root GPU self.root_gpu = gpu_nb @@ -989,13 +1004,7 @@ class Trainer(TrainerIOMixin): ref_model.trainer = self # set local properties on the model - ref_model.on_gpu = self.on_gpu - ref_model.single_gpu = self.single_gpu - ref_model.use_dp = self.use_dp - ref_model.use_ddp = self.use_ddp - ref_model.use_ddp2 = self.use_ddp2 - ref_model.use_amp = self.use_amp - ref_model.testing = self.testing + self.__copy_trainer_model_properties(ref_model) # link up experiment object if self.logger is not None: @@ -1199,8 +1208,8 @@ class Trainer(TrainerIOMixin): self.logger.save() def test(self, model=None): + self.testing = True if model is not None: - self.testing = True self.fit(model) else: self.__run_evaluation(test=True) @@ -1305,7 +1314,7 @@ class Trainer(TrainerIOMixin): if k not in ['progress_bar', 'log']: callback_metrics[k] = v - if train and self.use_dp or self.use_ddp2: + if train and (self.use_dp or self.use_ddp2): nb_gpus = self.num_gpus callback_metrics = reduce_distributed_output(callback_metrics, nb_gpus) @@ -1319,7 +1328,7 @@ class Trainer(TrainerIOMixin): progress_output = output['progress_bar'] # reduce progress metrics for tqdm when using dp - if train and self.use_dp or self.use_ddp2: + if train and (self.use_dp or self.use_ddp2): nb_gpus = self.num_gpus progress_output = reduce_distributed_output(progress_output, nb_gpus) @@ -1517,13 +1526,13 @@ class Trainer(TrainerIOMixin): model.on_pre_performance_check() # select dataloaders - dataloaders = self.get_val_dataloaders() - max_batches = self.nb_val_batches - - # calculate max batches to use if test: dataloaders = self.get_test_dataloaders() max_batches = self.nb_test_batches + else: + # val + dataloaders = self.get_val_dataloaders() + max_batches = self.nb_val_batches # cap max batches to 1 when using fast_dev_run if self.fast_dev_run: