From 8d564b5e38d1a1f820304a27f2d615d8bd4f401d Mon Sep 17 00:00:00 2001 From: Peter Yu <2057325+yukw777@users.noreply.github.com> Date: Thu, 30 Apr 2020 07:57:24 -0400 Subject: [PATCH] call on_load_checkpoint() when resuming from checkpoint (#1666) --- CHANGELOG.md | 1 + pytorch_lightning/trainer/training_io.py | 4 ++++ tests/trainer/test_trainer.py | 15 +++++++++++---- 3 files changed, 16 insertions(+), 4 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 85edc738..10ec061f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,6 +18,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed broken link in PR template ([#1675](https://github.com/PyTorchLightning/pytorch-lightning/pull/1675)) - Fixed ModelCheckpoint not None checking filepath ([1654](https://github.com/PyTorchLightning/pytorch-lightning/pull/1654)) +- Trainer now calls `on_load_checkpoint()` when resuming from a checkpoint ([1666](https://github.com/PyTorchLightning/pytorch-lightning/pull/1666)) ## [0.7.5] - 2020-04-27 diff --git a/pytorch_lightning/trainer/training_io.py b/pytorch_lightning/trainer/training_io.py index 82bc0829..393d6540 100644 --- a/pytorch_lightning/trainer/training_io.py +++ b/pytorch_lightning/trainer/training_io.py @@ -278,6 +278,10 @@ class TrainerIOMixin(ABC): # load the state_dict on the model automatically model.load_state_dict(checkpoint['state_dict']) + + # give model a chance to load something + model.on_load_checkpoint(checkpoint) + if on_gpu: model.cuda(self.root_gpu) diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index 18cc2586..cb650fd8 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -309,8 +309,8 @@ def test_model_freeze_unfreeze(): model.unfreeze() -def test_resume_from_checkpoint_epoch_restored(tmpdir): - """Verify resuming from checkpoint runs the right number of epochs""" +def test_resume_from_checkpoint(tmpdir): + """Verify resuming from checkpoint (epoch, batch numbers and on_load_checkpoint())""" import types tutils.reset_seed() @@ -322,6 +322,7 @@ def test_resume_from_checkpoint_epoch_restored(tmpdir): model = LightningTestModel(hparams) model.num_epochs_seen = 0 model.num_batches_seen = 0 + model.num_on_load_checkpoint_called = 0 def increment_epoch(self): self.num_epochs_seen += 1 @@ -329,10 +330,14 @@ def test_resume_from_checkpoint_epoch_restored(tmpdir): def increment_batch(self, _): self.num_batches_seen += 1 - # Bind the increment_epoch function on_epoch_end so that the - # model keeps track of the number of epochs it has seen. + def increment_on_load_checkpoint(self, _): + self.num_on_load_checkpoint_called += 1 + + # Bind methods to keep track of epoch numbers, batch numbers it has seen + # as well as number of times it has called on_load_checkpoint() model.on_epoch_end = types.MethodType(increment_epoch, model) model.on_batch_start = types.MethodType(increment_batch, model) + model.on_load_checkpoint = types.MethodType(increment_on_load_checkpoint, model) return model model = _new_model() @@ -356,6 +361,7 @@ def test_resume_from_checkpoint_epoch_restored(tmpdir): assert model.num_epochs_seen == 2 assert model.num_batches_seen == training_batches * 2 + assert model.num_on_load_checkpoint_called == 0 # Other checkpoints can be uncommented if/when resuming mid-epoch is supported checkpoints = sorted(glob.glob(os.path.join(trainer.checkpoint_callback.dirpath, '*.ckpt'))) @@ -369,6 +375,7 @@ def test_resume_from_checkpoint_epoch_restored(tmpdir): new_trainer = Trainer(**trainer_options, resume_from_checkpoint=check) new_trainer.fit(next_model) assert state['global_step'] + next_model.num_batches_seen == training_batches * trainer_options['max_epochs'] + assert next_model.num_on_load_checkpoint_called == 1 def _init_steps_model():