mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
Improved docs for callbacks (#1370)
* improved docs for callbacks * class references * make doctest pass * doctests * fix lines too long * fix line too long * fix permission error in doctest * Apply suggestions from code review Co-Authored-By: Jirka Borovec <Borda@users.noreply.github.com> * fix doctest * fix default Co-authored-by: Jirka Borovec <Borda@users.noreply.github.com>
This commit is contained in:
co-authored by
Jirka Borovec
parent
22bedf9b57
commit
1f2da71069
@@ -3,6 +3,7 @@ Model Checkpointing
|
||||
===================
|
||||
|
||||
Automatically save model checkpoints during training.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -26,18 +27,19 @@ class ModelCheckpoint(Callback):
|
||||
|
||||
Example::
|
||||
|
||||
# no path
|
||||
ModelCheckpoint()
|
||||
# saves like /my/path/epoch_0.ckpt
|
||||
# custom path
|
||||
# saves a file like: my/path/epoch_0.ckpt
|
||||
>>> checkpoint_callback = ModelCheckpoint('my/path/')
|
||||
|
||||
# save any arbitrary metrics like and val_loss, etc in name
|
||||
ModelCheckpoint(filepath='/my/path/{epoch}-{val_loss:.2f}-{other_metric:.2f}')
|
||||
# saves file like: /my/path/epoch=2-val_loss=0.2_other_metric=0.3.ckpt
|
||||
# save any arbitrary metrics like `val_loss`, etc. in name
|
||||
# saves a file like: my/path/epoch=2-val_loss=0.2_other_metric=0.3.ckpt
|
||||
>>> checkpoint_callback = ModelCheckpoint(
|
||||
... filepath='my/path/{epoch}-{val_loss:.2f}-{other_metric:.2f}'
|
||||
... )
|
||||
|
||||
|
||||
monitor (str): quantity to monitor.
|
||||
verbose (bool): verbosity mode, False or True.
|
||||
save_top_k (int): if `save_top_k == k`,
|
||||
monitor: quantity to monitor.
|
||||
verbose: verbosity mode. Default: ``False``.
|
||||
save_top_k: if `save_top_k == k`,
|
||||
the best k models according to
|
||||
the quantity monitored will be saved.
|
||||
if ``save_top_k == 0``, no models are saved.
|
||||
@@ -46,7 +48,7 @@ class ModelCheckpoint(Callback):
|
||||
if ``save_top_k >= 2`` and the callback is called multiple
|
||||
times inside an epoch, the name of the saved file will be
|
||||
appended with a version count starting with `v0`.
|
||||
mode (str): one of {auto, min, max}.
|
||||
mode: one of {auto, min, max}.
|
||||
If ``save_top_k != 0``, the decision
|
||||
to overwrite the current save file is made
|
||||
based on either the maximization or the
|
||||
@@ -54,26 +56,29 @@ class ModelCheckpoint(Callback):
|
||||
this should be `max`, for `val_loss` this should
|
||||
be `min`, etc. In `auto` mode, the direction is
|
||||
automatically inferred from the name of the monitored quantity.
|
||||
save_weights_only (bool): if True, then only the model's weights will be
|
||||
saved (`model.save_weights(filepath)`), else the full model
|
||||
is saved (`model.save(filepath)`).
|
||||
period (int): Interval (number of epochs) between checkpoints.
|
||||
save_weights_only: if ``True``, then only the model's weights will be
|
||||
saved (``model.save_weights(filepath)``), else the full model
|
||||
is saved (``model.save(filepath)``).
|
||||
period: Interval (number of epochs) between checkpoints.
|
||||
|
||||
Example::
|
||||
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.callbacks import ModelCheckpoint
|
||||
>>> from pytorch_lightning import Trainer
|
||||
>>> from pytorch_lightning.callbacks import ModelCheckpoint
|
||||
|
||||
# saves checkpoints to my_path whenever 'val_loss' has a new min
|
||||
checkpoint_callback = ModelCheckpoint(filepath='my_path')
|
||||
Trainer(checkpoint_callback=checkpoint_callback)
|
||||
# saves checkpoints to 'my/path/' whenever 'val_loss' has a new min
|
||||
>>> checkpoint_callback = ModelCheckpoint(filepath='my/path/')
|
||||
>>> trainer = Trainer(checkpoint_callback=checkpoint_callback)
|
||||
|
||||
# save epoch and val_loss in name
|
||||
ModelCheckpoint(filepath='/my/path/here/sample-mnist_{epoch:02d}-{val_loss:.2f}')
|
||||
# saves file like: /my/path/here/sample-mnist_epoch=02_val_loss=0.32.ckpt
|
||||
# saves a file like: my/path/sample-mnist_epoch=02_val_loss=0.32.ckpt
|
||||
>>> checkpoint_callback = ModelCheckpoint(
|
||||
... filepath='my/path/sample-mnist_{epoch:02d}-{val_loss:.2f}'
|
||||
... )
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, filepath, monitor: str = 'val_loss', verbose: bool = False,
|
||||
def __init__(self, filepath: str, monitor: str = 'val_loss', verbose: bool = False,
|
||||
save_top_k: int = 1, save_weights_only: bool = False,
|
||||
mode: str = 'auto', period: int = 1, prefix: str = ''):
|
||||
super().__init__()
|
||||
@@ -137,9 +142,10 @@ class ModelCheckpoint(Callback):
|
||||
return self.monitor_op(current, self.best_k_models[self.kth_best_model])
|
||||
|
||||
def format_checkpoint_name(self, epoch, metrics, ver=None):
|
||||
"""Generate a filename according define template.
|
||||
"""Generate a filename according to the defined template.
|
||||
|
||||
Example::
|
||||
|
||||
Examples:
|
||||
>>> tmpdir = os.path.dirname(__file__)
|
||||
>>> ckpt = ModelCheckpoint(os.path.join(tmpdir, '{epoch}'))
|
||||
>>> os.path.basename(ckpt.format_checkpoint_name(0, {}))
|
||||
|
||||
Reference in New Issue
Block a user