mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
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:
+139
-26
@@ -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__])
|
||||
|
||||
Reference in New Issue
Block a user