diff --git a/src/pytorch-lightning/models/trainer.py b/src/pytorch-lightning/models/trainer.py index 7d6f43a8..c1ba181f 100644 --- a/src/pytorch-lightning/models/trainer.py +++ b/src/pytorch-lightning/models/trainer.py @@ -267,8 +267,7 @@ class Trainer(TrainerIO): break # give model a chance to end epoch early - if self.model.should_stop_epoch: - self.model.should_stop_epoch = False + if self.model.should_stop_epoch(data_batch): break # --------------- diff --git a/src/pytorch-lightning/root_module/hooks.py b/src/pytorch-lightning/root_module/hooks.py index 0d903f66..02471b29 100644 --- a/src/pytorch-lightning/root_module/hooks.py +++ b/src/pytorch-lightning/root_module/hooks.py @@ -19,3 +19,5 @@ class ModelHooks(torch.nn.Module): def on_post_performance_check(self): pass + def should_stop_epoch(self, data_batch): + return False diff --git a/src/pytorch-lightning/root_module/root_module.py b/src/pytorch-lightning/root_module/root_module.py index b12a2087..ab49a9a1 100644 --- a/src/pytorch-lightning/root_module/root_module.py +++ b/src/pytorch-lightning/root_module/root_module.py @@ -24,7 +24,6 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks): self.overfit = hparams.overfit self.gradient_clip = hparams.gradient_clip self.num = 2 - self.should_stop_epoch = False # track if gpu was requested for checkpointing self.on_gpu = False