espnet3.components.callbacks.default_callbacks.get_default_callbacks
Less than 1 minute
espnet3.components.callbacks.default_callbacks.get_default_callbacks
espnet3.components.callbacks.default_callbacks.get_default_callbacks(exp_dir: str = './exp', log_interval: int = 500, best_model_criterion: List[Tuple[str, int, str]] | List[List] = [('valid/loss', 3, 'min')]) → List[Callback]
Return a list of callbacks tailored for most training workflows.
Includes.
ModelCheckpointfor saving the last model checkpoint (save_last)- One or more
ModelCheckpointinstances for saving the top-K checkpoints according to : specific metricsAverageCheckpointsCallbackto compute and save the average model from top-K : checkpointsLearningRateMonitorto track and log learning rates during trainingMetricsLoggerto emit train and validation summariesTQDMProgressBarto show a rich progress bar during training
- Parameters:
- exp_dir (str) – Directory to store checkpoints and logs.
- log_interval (int) – Frequency (in training steps) to refresh the progress bar.
- best_model_criterion (List *[*Tuple *[*str , int , str ] ]) – Criteria for saving top-K checkpoints. Each tuple is
(name, top_k, mode), wherenameis the validation value to monitor,top_kis the number of retained checkpoints, andmodeis"min"or"max".
- Returns: A list of callbacks to be passed to the PyTorch Lightning : Trainer.
- Return type: List[Callback]
Example
>>> from default_callbacks import get_default_callbacks
>>> callbacks = get_default_callbacks(
... exp_dir="./exp",
... log_interval=100,
... best_model_criterion=[("val/loss", 5, "min"), ("val/acc", 3, "max")]
... )
>>> trainer = Trainer(callbacks=callbacks, ...)