From 046082139876562c0843ac0aaf178b305fb4a9ac Mon Sep 17 00:00:00 2001 From: William Falcon Date: Tue, 25 Jun 2019 20:12:41 -0400 Subject: [PATCH] updated args --- docs/source/examples/example_model.py | 4 ++-- pytorch_lightning/pt_overrides/override_data_parallel.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/docs/source/examples/example_model.py b/docs/source/examples/example_model.py index c3d64e67..a8954d90 100644 --- a/docs/source/examples/example_model.py +++ b/docs/source/examples/example_model.py @@ -78,7 +78,7 @@ class ExampleModel(RootModule): output = OrderedDict({ 'loss_val': loss_val, }) - return torch.tensor(4) + return output def validation_step(self, data_batch, batch_i): """ @@ -102,7 +102,7 @@ class ExampleModel(RootModule): 'loss_val': loss_val, 'val_acc': val_acc, }) - return torch.tensor(4) + return output def validation_end(self, outputs): diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index 1076c8a6..9be6dab1 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -72,9 +72,9 @@ def parallel_apply(modules, inputs, kwargs_tup=None, devices=None): # --------------- # CHANGE if module.training: - return module.training_step(*input, **kwargs) + output = module.training_step(*input, **kwargs) else: - return module.validation_step(*input, **kwargs) + output = module.validation_step(*input, **kwargs) # --------------- with lock: