diff --git a/examples/new_project_templates/lightning_module_template.py b/examples/new_project_templates/lightning_module_template.py index 91fe8c3e..a0c65af5 100644 --- a/examples/new_project_templates/lightning_module_template.py +++ b/examples/new_project_templates/lightning_module_template.py @@ -116,7 +116,7 @@ class LightningTemplateModel(LightningModule): :param outputs: list of individual outputs of each validation step :return: """ - return outputs.mean() + return torch.stack(outputs).mean() val_loss_mean = 0 val_acc_mean = 0