mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
updated docs
This commit is contained in:
@@ -68,6 +68,10 @@ class CoolModel(ptl.LightningModule):
|
|||||||
y_hat = self.forward(x)
|
y_hat = self.forward(x)
|
||||||
return {'val_loss': self.my_loss(y_hat, y)}
|
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):
|
def configure_optimizers(self):
|
||||||
return [torch.optim.Adam(self.parameters(), lr=0.02)]
|
return [torch.optim.Adam(self.parameters(), lr=0.02)]
|
||||||
|
|
||||||
|
|||||||
@@ -57,6 +57,10 @@ class CoolModel(ptl.LightningModule):
|
|||||||
y_hat = self.forward(x)
|
y_hat = self.forward(x)
|
||||||
return {'val_loss': self.my_loss(y_hat, y)}
|
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):
|
def configure_optimizers(self):
|
||||||
return [torch.optim.Adam(self.parameters(), lr=0.02)]
|
return [torch.optim.Adam(self.parameters(), lr=0.02)]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user