From 8c4c7b105e16fbe255e4715f54af2fa5d2a12fad Mon Sep 17 00:00:00 2001 From: Fabio Natanael Kepler Date: Sun, 17 May 2020 14:24:17 +0100 Subject: [PATCH] Fix `save_weights_only` flag in ModelCheckpoint (#1780) * Add flag to `dump_checkpoint` for only including weights `ModelCheckpoint` then passes `self.save_weights_only` to the save function. * Fix tests and add changelog entry * Add check and descriptive message when training state is restored from a weights only checkpoint Also add a test for making sure `ModelCheckpoint.save_weights_only` works as expected. * Fix weights-only test to properly match expected exception * Apply suggestions from code review Co-authored-by: Jirka Borovec Co-authored-by: Jirka Borovec --- CHANGELOG.md | 2 + .../callbacks/model_checkpoint.py | 2 +- pytorch_lightning/trainer/training_io.py | 51 +++++++++++-------- tests/trainer/test_trainer.py | 40 ++++++++++++++- 4 files changed, 71 insertions(+), 24 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 24bd02fb..fa568479 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -71,6 +71,8 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed an issue with Trainer constructor silently ignoring unkown/misspelled arguments ([#1820](https://github.com/PyTorchLightning/pytorch-lightning/pull/1820)) +- Fixed `save_weights_only` in ModelCheckpoint ([#1780](https://github.com/PyTorchLightning/pytorch-lightning/pull/1780)) + ## [0.7.5] - 2020-04-27 ### Changed diff --git a/pytorch_lightning/callbacks/model_checkpoint.py b/pytorch_lightning/callbacks/model_checkpoint.py index ec403ee6..a65855e4 100644 --- a/pytorch_lightning/callbacks/model_checkpoint.py +++ b/pytorch_lightning/callbacks/model_checkpoint.py @@ -139,7 +139,7 @@ class ModelCheckpoint(Callback): # delegate the saving to the model if self.save_function is not None: - self.save_function(filepath) + self.save_function(filepath, self.save_weights_only) else: raise ValueError(".save_function() not set") diff --git a/pytorch_lightning/trainer/training_io.py b/pytorch_lightning/trainer/training_io.py index 4f83949f..11771b21 100644 --- a/pytorch_lightning/trainer/training_io.py +++ b/pytorch_lightning/trainer/training_io.py @@ -256,8 +256,8 @@ class TrainerIOMixin(ABC): torch.save(checkpoint, tmp_path) os.replace(tmp_path, filepath) - def save_checkpoint(self, filepath): - checkpoint = self.dump_checkpoint() + def save_checkpoint(self, filepath, weights_only: bool = False): + checkpoint = self.dump_checkpoint(weights_only) if self.proc_rank == 0: # do the actual save @@ -306,42 +306,43 @@ class TrainerIOMixin(ABC): # load training state (affects trainer only) self.restore_training_state(checkpoint) - def dump_checkpoint(self): + def dump_checkpoint(self, weights_only: bool = False): checkpoint = { 'epoch': self.current_epoch + 1, 'global_step': self.global_step + 1, } - if self.checkpoint_callback is not None and self.checkpoint_callback is not False: - checkpoint['checkpoint_callback_best'] = self.checkpoint_callback.best + if not weights_only: + if self.checkpoint_callback: + checkpoint['checkpoint_callback_best'] = self.checkpoint_callback.best - if self.early_stop_callback is not None and self.checkpoint_callback is not False: - checkpoint['early_stop_callback_wait'] = self.early_stop_callback.wait - checkpoint['early_stop_callback_patience'] = self.early_stop_callback.patience + if self.early_stop_callback: + checkpoint['early_stop_callback_wait'] = self.early_stop_callback.wait + checkpoint['early_stop_callback_patience'] = self.early_stop_callback.patience - # save optimizers - optimizer_states = [] - for i, optimizer in enumerate(self.optimizers): - optimizer_states.append(optimizer.state_dict()) + # save optimizers + optimizer_states = [] + for i, optimizer in enumerate(self.optimizers): + optimizer_states.append(optimizer.state_dict()) - checkpoint['optimizer_states'] = optimizer_states + checkpoint['optimizer_states'] = optimizer_states - # save lr schedulers - lr_schedulers = [] - for scheduler in self.lr_schedulers: - lr_schedulers.append(scheduler['scheduler'].state_dict()) + # save lr schedulers + lr_schedulers = [] + for scheduler in self.lr_schedulers: + lr_schedulers.append(scheduler['scheduler'].state_dict()) - checkpoint['lr_schedulers'] = lr_schedulers + checkpoint['lr_schedulers'] = lr_schedulers + + # save native amp scaling + if self.use_amp and self.use_native_amp: + checkpoint['native_amp_scaling_state'] = self.scaler.state_dict() # add the hparams and state_dict from the model model = self.get_model() checkpoint['state_dict'] = model.state_dict() - # save native amp scaling - if self.use_amp and self.use_native_amp: - checkpoint['native_amp_scaling_state'] = self.scaler.state_dict() - if hasattr(model, "hparams") and model.hparams is not None: parsing.clean_namespace(model.hparams) if isinstance(model.hparams, dict): @@ -391,6 +392,12 @@ class TrainerIOMixin(ABC): :param checkpoint: :return: """ + if 'optimizer_states' not in checkpoint or 'lr_schedulers' not in checkpoint: + raise KeyError( + 'Trying to restore training state but checkpoint contains only the model.' + ' This is probably due to `ModelCheckpoint.save_weights_only` being set to `True`.' + ) + if self.checkpoint_callback is not None and self.checkpoint_callback is not False: self.checkpoint_callback.best = checkpoint['checkpoint_callback_best'] diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index 0116125a..1c2c1691 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -270,7 +270,7 @@ def test_dp_output_reduce(): def test_model_checkpoint_options(tmpdir, save_top_k, file_prefix, expected_files): """Test ModelCheckpoint options.""" - def mock_save_function(filepath): + def mock_save_function(filepath, *args): open(filepath, 'a').close() # simulated losses @@ -296,6 +296,44 @@ def test_model_checkpoint_options(tmpdir, save_top_k, file_prefix, expected_file assert fname in file_lists +def test_model_checkpoint_only_weights(tmpdir): + """Tests use case where ModelCheckpoint is configured to save only model weights, and + user tries to load checkpoint to resume training. + """ + model = EvalModelTemplate() + + trainer = Trainer( + max_epochs=1, + checkpoint_callback=ModelCheckpoint(tmpdir, save_weights_only=True) + ) + # fit model + result = trainer.fit(model) + # training complete + assert result == 1, 'training failed to complete' + + checkpoint_path = list(trainer.checkpoint_callback.best_k_models.keys())[0] + + # assert saved checkpoint has no trainer data + checkpoint = torch.load(checkpoint_path) + assert 'optimizer_states' not in checkpoint, 'checkpoint should contain only model weights' + assert 'lr_schedulers' not in checkpoint, 'checkpoint should contain only model weights' + + # assert loading model works when checkpoint has only weights + assert EvalModelTemplate.load_from_checkpoint(checkpoint_path=checkpoint_path) + + # directly save model + new_weights_path = os.path.join(tmpdir, 'save_test.ckpt') + trainer.save_checkpoint(new_weights_path, weights_only=True) + # assert saved checkpoint has no trainer data + checkpoint = torch.load(new_weights_path) + assert 'optimizer_states' not in checkpoint, 'checkpoint should contain only model weights' + assert 'lr_schedulers' not in checkpoint, 'checkpoint should contain only model weights' + + # assert restoring train state fails + with pytest.raises(KeyError, match='checkpoint contains only the model'): + trainer.restore_training_state(checkpoint) + + def test_model_freeze_unfreeze(): model = EvalModelTemplate()