From 9340e0a09167cfbd6f832b719a0d80109d865e2b Mon Sep 17 00:00:00 2001 From: William Falcon Date: Wed, 3 Jul 2019 16:51:32 -0400 Subject: [PATCH] clean up dead code --- pytorch_lightning/models/trainer.py | 20 ++++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 254db80b..3e6dc09c 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -209,11 +209,11 @@ class Trainer(TrainerIO): # ----------------- # RUN VALIDATION STEP # ----------------- - # if self.data_parallel: - # output = model(data_batch, batch_i) - # output = reduce_distributed_output(output, len(self.data_parallel_device_ids)) - # else: - output = model.validation_step(data_batch, batch_i) + if self.data_parallel: + output = model(data_batch, batch_i) + # output = reduce_distributed_output(output, len(self.data_parallel_device_ids)) + else: + output = model.validation_step(data_batch, batch_i) outputs.append(output) @@ -483,11 +483,11 @@ class Trainer(TrainerIO): # forward pass # return a scalar value and a dic with tqdm metrics - # if self.data_parallel: - # output = self.model(data_batch, batch_nb) - # output = reduce_distributed_output(output, len(self.data_parallel_device_ids)) - # else: - output = self.model.training_step(data_batch, batch_nb) + if self.data_parallel: + output = self.model(data_batch, batch_nb) + # output = reduce_distributed_output(output, len(self.data_parallel_device_ids)) + else: + output = self.model.training_step(data_batch, batch_nb) model_specific_tqdm_metrics_dic = output['tqdm_metrics'] loss = output['loss']