Callbacks [wip] (#889)

* Add callback system + associated test

* Add trainer and pl_module args to callback methods

* typing

* typo in docstring

* Switch to on_.*_start()

* fix on_test_start

* fix the mess after rebasing
This commit is contained in:
Hadrien Mary
2020-02-25 23:17:27 -05:00
committed by GitHub
parent 96b058c5fa
commit be244560b2
14 changed files with 407 additions and 87 deletions
+139 -26
View File
@@ -23,10 +23,13 @@ from tests.models import (
LightValStepFitSingleDataloaderMixin,
LightTrainDataloader,
LightTestDataloader,
LightValidationMixin,
LightTestMixin
)
from pytorch_lightning.core.lightning import load_hparams_from_tags_csv
from pytorch_lightning.trainer.logging import TrainerLoggingMixin
from pytorch_lightning.utilities.debugging import MisconfigurationException
from pytorch_lightning import Callback
def test_no_val_module(tmpdir):
@@ -242,13 +245,12 @@ def test_model_checkpoint_options(tmp_path):
checkpoint_callback = ModelCheckpoint(save_dir, save_top_k=-1, verbose=1)
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for i, loss in enumerate(losses):
checkpoint_callback._trainer.current_epoch = i
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = i
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = set(os.listdir(save_dir))
@@ -266,13 +268,12 @@ def test_model_checkpoint_options(tmp_path):
checkpoint_callback = ModelCheckpoint(save_dir, save_top_k=0, verbose=1)
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for i, loss in enumerate(losses):
checkpoint_callback._trainer.current_epoch = i
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = i
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = os.listdir(save_dir)
@@ -286,13 +287,12 @@ def test_model_checkpoint_options(tmp_path):
checkpoint_callback = ModelCheckpoint(save_dir, save_top_k=1, verbose=1, prefix='test_prefix')
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for i, loss in enumerate(losses):
checkpoint_callback._trainer.current_epoch = i
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = i
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = set(os.listdir(save_dir))
@@ -310,13 +310,12 @@ def test_model_checkpoint_options(tmp_path):
open(f'{save_dir}/other_file.ckpt', 'a').close()
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for i, loss in enumerate(losses):
checkpoint_callback._trainer.current_epoch = i
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = i
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = set(os.listdir(save_dir))
@@ -335,13 +334,12 @@ def test_model_checkpoint_options(tmp_path):
checkpoint_callback = ModelCheckpoint(save_dir, save_top_k=4, verbose=1)
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for loss in losses:
checkpoint_callback._trainer.current_epoch = 0
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = 0
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = set(os.listdir(save_dir))
@@ -357,13 +355,12 @@ def test_model_checkpoint_options(tmp_path):
checkpoint_callback = ModelCheckpoint(save_dir, save_top_k=3, verbose=1)
checkpoint_callback.save_function = mock_save_function
trainer = Trainer()
checkpoint_callback.set_trainer(trainer)
# emulate callback's calls during the training
for loss in losses:
checkpoint_callback._trainer.current_epoch = 0
checkpoint_callback._trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end()
trainer.current_epoch = 0
trainer.callback_metrics = {'val_loss': loss}
checkpoint_callback.on_validation_end(trainer, trainer.get_model())
file_lists = set(os.listdir(save_dir))
@@ -798,8 +795,9 @@ def test_benchmark_option(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValidationMultipleDataloadersMixin,
LightningTestModelBase
LightValidationMultipleDataloadersMixin,
LightTrainDataloader,
TestModelBase
):
pass
@@ -858,5 +856,120 @@ def test_testpass_overrides(tmpdir):
Trainer().test(model)
def test_trainer_callback_system(tmpdir):
"""Test the callback system."""
class CurrentTestModel(
LightTrainDataloader,
LightTestMixin,
LightValidationMixin,
TestModelBase,
):
pass
hparams = tutils.get_hparams()
model = CurrentTestModel(hparams)
class TestCallback(Callback):
def __init__(self):
super().__init__()
self.on_init_start_called = False
self.on_init_end_called = False
self.on_fit_start_called = False
self.on_fit_end_called = False
self.on_epoch_start_called = False
self.on_epoch_end_called = False
self.on_batch_start_called = False
self.on_batch_end_called = False
self.on_train_start_called = False
self.on_train_end_called = False
self.on_validation_start_called = False
self.on_validation_end_called = False
self.on_test_start_called = False
self.on_test_end_called = False
def on_init_start(self, trainer, pl_module):
self.on_init_start_called = True
def on_init_end(self, trainer, pl_module):
self.on_init_end_called = True
def on_fit_start(self, trainer, pl_module):
self.on_fit_start_called = True
def on_fit_end(self, trainer, pl_module):
self.on_fit_end_called = True
def on_epoch_start(self, trainer, pl_module):
self.on_epoch_start_called = True
def on_epoch_end(self, trainer, pl_module):
self.on_epoch_end_called = True
def on_batch_start(self, trainer, pl_module):
self.on_batch_start_called = True
def on_batch_end(self, trainer, pl_module):
self.on_batch_end_called = True
def on_train_start(self, trainer, pl_module):
self.on_train_start_called = True
def on_train_end(self, trainer, pl_module):
self.on_train_end_called = True
def on_validation_start(self, trainer, pl_module):
self.on_validation_start_called = True
def on_validation_end(self, trainer, pl_module):
self.on_validation_end_called = True
def on_test_start(self, trainer, pl_module):
self.on_test_start_called = True
def on_test_end(self, trainer, pl_module):
self.on_test_end_called = True
test_callback = TestCallback()
trainer_options = {}
trainer_options['callbacks'] = [test_callback]
trainer_options['max_epochs'] = 1
trainer_options['val_percent_check'] = 0.1
trainer_options['train_percent_check'] = 0.2
trainer_options['show_progress_bar'] = False
assert not test_callback.on_init_start_called
assert not test_callback.on_init_end_called
# fit model
trainer = Trainer(**trainer_options)
assert trainer.callbacks[0] == test_callback
assert test_callback.on_init_start_called
assert test_callback.on_init_end_called
assert not test_callback.on_fit_start_called
assert not test_callback.on_fit_start_called
trainer.fit(model)
assert test_callback.on_fit_start_called
assert test_callback.on_fit_end_called
assert test_callback.on_epoch_start_called
assert test_callback.on_epoch_start_called
assert test_callback.on_batch_start_called
assert test_callback.on_batch_end_called
assert test_callback.on_train_start_called
assert test_callback.on_train_end_called
assert test_callback.on_validation_start_called
assert test_callback.on_validation_end_called
assert not test_callback.on_test_start_called
assert not test_callback.on_test_end_called
trainer.test()
assert test_callback.on_test_start_called
assert test_callback.on_test_end_called
# if __name__ == '__main__':
# pytest.main([__file__])