diff --git a/examples/new_project_templates/lightning_module_template.py b/examples/new_project_templates/lightning_module_template.py index 62a34c9a..94e3407d 100644 --- a/examples/new_project_templates/lightning_module_template.py +++ b/examples/new_project_templates/lightning_module_template.py @@ -105,7 +105,7 @@ class LightningTemplateModel(LightningModule): # can also return just a scalar instead of a dict (return loss_val) return output - def validation_step(self, data_batch, batch_i, dataloader_i): + def validation_step(self, data_batch, batch_i): """ Lightning calls this inside the validation loop :param data_batch: @@ -218,7 +218,7 @@ class LightningTemplateModel(LightningModule): @pl.data_loader def val_dataloader(self): print('val data loader called') - return [self.__dataloader(train=False) for i in range(2)] + return self.__dataloader(train=False) @pl.data_loader def test_dataloader(self): diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 00e6dece..4f58cf7a 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -377,6 +377,32 @@ class Trainer(TrainerIO): self.tqdm_metrics[k] = v + def __validation_forward(self, model, data_batch, batch_i, dataloader_i): + # make dataloader_i arg in validation_step optional + args = [data_batch, batch_i] + if len(self.val_dataloader) > 1: + args.append(dataloader_i) + + if self.use_ddp: + output = model(*args) + elif self.use_dp: + output = model(*args) + elif self.single_gpu: + # put inputs on gpu manually + gpu_id = self.data_parallel_device_ids[0] + for i, x in enumerate(data_batch): + if isinstance(x, torch.Tensor): + data_batch[i] = x.cuda(gpu_id) + + # do non dp, ddp step + output = model.validation_step(*args) + + else: + # CPU + output = model.validation_step(*args) + + return output + def validate(self, model, dataloader, max_batches, dataloader_i): """ Run validation code @@ -409,23 +435,9 @@ class Trainer(TrainerIO): # ----------------- # RUN VALIDATION STEP # ----------------- - if self.use_ddp: - output = model(data_batch, batch_i, dataloader_i) - elif self.use_dp: - output = model(data_batch, batch_i, dataloader_i) - elif self.single_gpu: - # put inputs on gpu manually - gpu_id = self.data_parallel_device_ids[0] - for i, x in enumerate(data_batch): - if isinstance(x, torch.Tensor): - data_batch[i] = x.cuda(gpu_id) - - # do non dp, ddp step - output = model.validation_step(data_batch, batch_i, dataloader_i) - - else: - output = model.validation_step(data_batch, batch_i, dataloader_i) + output = self.__validation_forward(model, data_batch, batch_i, dataloader_i) + # track outputs for collation outputs.append(output) # batch done diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index a9ba6ebd..05f21817 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -33,11 +33,21 @@ class LightningModule(GradInformation, ModelIO, ModelHooks): """ raise NotImplementedError - def validation_step(self, data_batch, batch_nb): + def training_step(self, *args, **kwargs): + """ + return loss, dict with metrics for tqdm + :param called with batch, batch_nb + additional: optimizer_i if multiple optimizers used + :return: + """ + raise NotImplementedError + + def validation_step(self, *args, **kwargs): """ return whatever outputs will need to be aggregated in validation_end OPTIONAL - :param data_batch: + :param called with batch, batch_nb + additional: dataset_i if multiple val datasets used :return: """ pass @@ -51,14 +61,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks): """ pass - def training_step(self, data_batch, batch_nb): - """ - return loss, dict with metrics for tqdm - :param data_batch: - :return: - """ - raise NotImplementedError - def configure_optimizers(self): """ Return a list of optimizers and a list of schedulers (could be empty) diff --git a/pytorch_lightning/testing/lm_test_module.py b/pytorch_lightning/testing/lm_test_module.py index e64b3ef5..ce0e7c3a 100644 --- a/pytorch_lightning/testing/lm_test_module.py +++ b/pytorch_lightning/testing/lm_test_module.py @@ -231,7 +231,7 @@ class LightningTestModel(LightningModule): @data_loader def val_dataloader(self): - return self.__dataloader(train=False) + return [self.__dataloader(train=False), self.__dataloader(train=False)] @data_loader def test_dataloader(self): diff --git a/pytorch_lightning/testing/no_val_end_module.py b/pytorch_lightning/testing/no_val_end_module.py index 8c21671e..897abb05 100644 --- a/pytorch_lightning/testing/no_val_end_module.py +++ b/pytorch_lightning/testing/no_val_end_module.py @@ -109,7 +109,7 @@ class NoValEndTestModel(LightningModule): if self.trainer.batch_nb % 2 == 0: return loss_val - def validation_step(self, data_batch, batch_i, dataloader_i): + def validation_step(self, data_batch, batch_nb): """ Lightning calls this inside the validation loop :param data_batch: @@ -135,16 +135,16 @@ class NoValEndTestModel(LightningModule): val_acc = val_acc.unsqueeze(0) # alternate possible outputs to test - if batch_i % 1 == 0: + if batch_nb % 1 == 0: output = OrderedDict({ 'val_loss': loss_val, 'val_acc': val_acc, }) return output - if batch_i % 2 == 0: + if batch_nb % 2 == 0: return val_acc - if batch_i % 3 == 0: + if batch_nb % 3 == 0: output = OrderedDict({ 'val_loss': loss_val, 'val_acc': val_acc, diff --git a/tests/test_models.py b/tests/test_models.py index 66def444..992bc7a8 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -767,7 +767,7 @@ def test_multiple_val_dataloader(): :return: """ hparams = get_hparams() - model = LightningTemplateModel(hparams) + model = LightningTestModel(hparams) save_dir = init_save_dir()