From 3dd0b8c186a2f7a15790b65d16ac73dd49cccbda Mon Sep 17 00:00:00 2001 From: Jeremy Jordan <13970565+jeremyjordan@users.noreply.github.com> Date: Sat, 14 Dec 2019 23:23:44 -0500 Subject: [PATCH] fix metric name to work with default earlystopping (#628) --- pytorch_lightning/core/__init__.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/pytorch_lightning/core/__init__.py b/pytorch_lightning/core/__init__.py index 0bb2323f..c2694eab 100644 --- a/pytorch_lightning/core/__init__.py +++ b/pytorch_lightning/core/__init__.py @@ -48,8 +48,8 @@ Minimal example def validation_end(self, outputs): # OPTIONAL - avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean() - return {'avg_val_loss': avg_loss} + val_loss_mean = torch.stack([x['val_loss'] for x in outputs]).mean() + return {'val_loss': val_loss_mean} def test_step(self, batch, batch_idx): # OPTIONAL @@ -59,8 +59,8 @@ Minimal example def test_end(self, outputs): # OPTIONAL - avg_loss = torch.stack([x['test_loss'] for x in outputs]).mean() - return {'avg_test_loss': avg_loss} + test_loss_mean = torch.stack([x['test_loss'] for x in outputs]).mean() + return {'test_loss': test_loss_mean} def configure_optimizers(self): # REQUIRED