made load from checkpoint flexible (#1307)

* made load from checkpoint flexible

* made load from checkpoint flexible

* made load from checkpoint flexible
This commit is contained in:
William Falcon
2020-03-30 18:28:51 -04:00
committed by GitHub
parent 31017120fd
commit 31a658e558
3 changed files with 55 additions and 8 deletions
+1
View File
@@ -21,6 +21,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
### Changed
- Enhanced load_from_checkpoint to also forward params to the model ([#1307](https://github.com/PyTorchLightning/pytorch-lightning/pull/1307))
- Made `evalaute` method private >> `Trainer._evaluate(...)`. ([#1260](https://github.com/PyTorchLightning/pytorch-lightning/pull/1260))
### Deprecated
+38 -5
View File
@@ -84,7 +84,7 @@ To save your own checkpoint call:
Checkpoint Loading
------------------
To load a model along with its weights, biases and hyperparameters use following method:
To load a model along with its weights, biases and hyperparameters use following method.
.. code-block:: python
@@ -92,9 +92,42 @@ To load a model along with its weights, biases and hyperparameters use following
model.eval()
y_hat = model(x)
A LightningModule is no different than a nn.Module. This means you can load it and use it for
predictions as you would a nn.Module.
The above only works if you used `hparams` in your model definition
.. code-block:: python
class MyModel(pl.LightningModule):
def __init__(self, hparams):
self.hparams = hparams
self.l1 = nn.Linear(hparams.in_dim, hparams.out_dim)
But if you don't and instead pass individual parameters
.. code-block:: python
class MyModel(pl.LightningModule):
def __init__(self, in_dim, out_dim):
self.l1 = nn.Linear(in_dim, out_dim)
you can restore the model like this
.. code-block:: python
model = MyModel.load_from_checkpoint(PATH, in_dim=128, out_dim=10)
.. note:: To restore the trainer state as well use
:meth:`pytorch_lightning.trainer.trainer.Trainer.resume_from_checkpoint`.
Restoring Training State
------------------------
If you don't just want to load weights, but instead restore the full training,
do the following:
.. code-block:: python
model = LitModel()
trainer = Trainer(resume_from_checkpoint='some/path/to/my_checkpoint.ckpt')
# automatically restores model, epoch, step, LR schedulers, apex, etc...
trainer.fit(model)
+16 -3
View File
@@ -1324,6 +1324,7 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
checkpoint_path: str,
map_location: Optional[Union[Dict[str, str], str, torch.device, int, Callable]] = None,
tags_csv: Optional[str] = None,
*args, **kwargs
) -> 'LightningModule':
r"""
@@ -1346,6 +1347,7 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
Args:
checkpoint_path: Path to checkpoint.
model_args: Any keyword args needed to init the model.
map_location:
If your checkpoint saved a GPU model and you now load on CPUs
or a different number of GPUs, use this to map to the new setup.
@@ -1387,6 +1389,14 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
tags_csv='/path/to/hparams_file.csv'
)
# or load passing whatever args the model takes to load
MyLightningModule.load_from_checkpoint(
'path/to/checkpoint.ckpt',
learning_rate=0.1,
layers=2,
pretrained_model=some_model
)
# predict
pretrained_model.eval()
pretrained_model.freeze()
@@ -1403,11 +1413,11 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
hparams.__setattr__('on_gpu', False)
checkpoint['hparams'] = vars(hparams)
model = cls._load_model_state(checkpoint)
model = cls._load_model_state(checkpoint, *args, **kwargs)
return model
@classmethod
def _load_model_state(cls, checkpoint: Dict[str, Any]) -> 'LightningModule':
def _load_model_state(cls, checkpoint: Dict[str, Any], *args, **kwargs) -> 'LightningModule':
cls_takes_hparams = 'hparams' in inspect.signature(cls.__init__).parameters
ckpt_hparams = checkpoint.get('hparams')
@@ -1433,7 +1443,10 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
# load the state_dict on the model automatically
model_args = [hparams] if hparams else []
model = cls(*model_args)
if len(model_args) > 0:
model = cls(*model_args)
else:
model = cls(*args, **kwargs)
model.load_state_dict(checkpoint['state_dict'])
# give model a chance to load something