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 <Borda@users.noreply.github.com>

Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com>
This commit is contained in:
Fabio Natanael Kepler
2020-05-17 09:24:17 -04:00
committed by GitHub
co-authored by Jirka Borovec
parent 769a459d27
commit 8c4c7b105e
4 changed files with 71 additions and 24 deletions
+2
View File
@@ -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
@@ -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")
+29 -22
View File
@@ -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']
+39 -1
View File
@@ -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()