mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
fixed val_loss for early stopping
This commit is contained in:
@@ -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,))])
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user