From a258d3d31bcb34266919d673c4895555fbee75f8 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sun, 26 Apr 2020 12:27:19 -0400 Subject: [PATCH] fixed val_loss for early stopping --- docs/source/introduction_guide.rst | 4 ++-- docs/source/new-project.rst | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/source/introduction_guide.rst b/docs/source/introduction_guide.rst index bed1f477..a7a406bb 100644 --- a/docs/source/introduction_guide.rst +++ b/docs/source/introduction_guide.rst @@ -654,7 +654,7 @@ sample split in the `train_dataloader` method. def validation_epoch_end(self, outputs): avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean() tensorboard_logs = {'val_loss': avg_loss} - return {'avg_val_loss': avg_loss, 'log': tensorboard_logs} + return {'val_loss': avg_loss, 'log': tensorboard_logs} def val_dataloader(self): transform=transforms.Compose([transforms.ToTensor(), @@ -710,7 +710,7 @@ Just like the validation loop, we define exactly the same steps for testing: def test_epoch_end(self, outputs): avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean() tensorboard_logs = {'val_loss': avg_loss} - return {'avg_val_loss': avg_loss, 'log': tensorboard_logs} + return {'val_loss': avg_loss, 'log': tensorboard_logs} def test_dataloader(self): transform=transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))]) diff --git a/docs/source/new-project.rst b/docs/source/new-project.rst index e58083eb..7d81ba44 100644 --- a/docs/source/new-project.rst +++ b/docs/source/new-project.rst @@ -100,7 +100,7 @@ To also add a validation loop add the following functions def validation_epoch_end(self, outputs): avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean() tensorboard_logs = {'val_loss': avg_loss} - return {'avg_val_loss': avg_loss, 'log': tensorboard_logs + return {'val_loss': avg_loss, 'log': tensorboard_logs def val_dataloader(self): # TODO: do a real train/val split