Split callbacks (#849)

* add .vscode in .gitignore

* Split callbacks in individual files + add a  property to Callback for easy trainer instance access

* formatting

* Add a conda env file for quick and easy env setup to develop on PL

* Adress comments

* add fix to kth_best_model

* add some typing to callbacks

* fix typo

* add autopep8 config to pyproject.toml

* format again

* format

* fix toml

* fix toml again

* consistent max line length in all config files

* remove conda env file

* Update pytorch_lightning/callbacks/early_stopping.py

Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com>

* Update pytorch_lightning/callbacks/model_checkpoint.py

Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com>

* docstring

* Update pytorch_lightning/callbacks/model_checkpoint.py

Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com>

* Update pytorch_lightning/callbacks/model_checkpoint.py

Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com>

* fix logic error

* format

* simplify if/else

* format

* fix linting issue in changelog

* edit changelog about new callback mechanism

* fix remaining formating issue on CHANGELOG

* remove lambda function because it's compatible with pickle (used during ddp)

Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com>
This commit is contained in:
Hadrien Mary
2020-02-22 21:45:34 -05:00
committed by GitHub
co-authored by Jirka Borovec
parent da2f11a9c4
commit 89d5772f55
11 changed files with 616 additions and 436 deletions
+68
View File
@@ -0,0 +1,68 @@
"""
Callbacks
=========
Callbacks supported by Lightning
"""
import abc
_NO_TRAINER_ERROR_MSG = ".set_trainer() should be called after the callback initialization"
class Callback(abc.ABC):
"""Abstract base class used to build new callbacks."""
def __init__(self):
self._trainer = None
@property
def trainer(self):
assert self._trainer is not None, _NO_TRAINER_ERROR_MSG
return self._trainer
def set_trainer(self, trainer):
"""Make a link to the trainer, so different things like `trainer.current_epoch`,
`trainer.batch_idx`, `trainer.global_step` can be used."""
self._trainer = trainer
def on_epoch_begin(self):
"""Called when the epoch begins."""
pass
def on_epoch_end(self):
"""Called when the epoch ends."""
pass
def on_batch_begin(self):
"""Called when the training batch begins."""
pass
def on_batch_end(self):
"""Called when the training batch ends."""
pass
def on_train_begin(self):
"""Called when the train begins."""
pass
def on_train_end(self):
"""Called when the train ends."""
pass
def on_validation_begin(self):
"""Called when the validation loop begins."""
pass
def on_validation_end(self):
"""Called when the validation loop ends."""
pass
def on_test_begin(self):
"""Called when the test begins."""
pass
def on_test_end(self):
"""Called when the test ends."""
pass