From 45625804610ebbd272ea5b87ed22e5b056c33014 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Thu, 25 Jul 2019 11:58:06 -0400 Subject: [PATCH] updated docs --- README.md | 4 ++++ docs/LightningModule/RequiredTrainerInterface.md | 4 ++++ 2 files changed, 8 insertions(+) diff --git a/README.md b/README.md index 8dc89cd4..9c0404fa 100644 --- a/README.md +++ b/README.md @@ -68,6 +68,10 @@ class CoolModel(ptl.LightningModule): y_hat = self.forward(x) return {'val_loss': self.my_loss(y_hat, y)} + def validation_end(self, outputs): + avg_loss = torch.stack([x for x in outputs['val_loss']]).mean() + return avg_loss + def configure_optimizers(self): return [torch.optim.Adam(self.parameters(), lr=0.02)] diff --git a/docs/LightningModule/RequiredTrainerInterface.md b/docs/LightningModule/RequiredTrainerInterface.md index 5faa67a7..9cf9b411 100644 --- a/docs/LightningModule/RequiredTrainerInterface.md +++ b/docs/LightningModule/RequiredTrainerInterface.md @@ -57,6 +57,10 @@ class CoolModel(ptl.LightningModule): y_hat = self.forward(x) return {'val_loss': self.my_loss(y_hat, y)} + def validation_end(self, outputs): + avg_loss = torch.stack([x for x in outputs['val_loss']]).mean() + return avg_loss + def configure_optimizers(self): return [torch.optim.Adam(self.parameters(), lr=0.02)]