espnet3.components.callbacks.ema.EMACallback
espnet3.components.callbacks.ema.EMACallback
class espnet3.components.callbacks.ema.EMACallback(decay: float = 0.9999, **ema_kwargs)
Bases: Callback
ESPnet3’s EMA callback system.
- Updates once per true optimizer step (not per micro-step).
- Swaps EMA weights in for validation/test, restores afterward.
- Saves under ‘ema_model_state_dict’.
Configure the callback.
- Parameters:
- decay – EMA decay rate, forwarded to
EMAasbeta. - **ema_kwargs – Extra keyword arguments forwarded to
EMA(espnet3.components.callbacks.vendored_ema.EMA). Note thatmodel,beta, andinclude_online_modelare already set by this callback. The remaining parameters:ema_model(default:None),update_after_step(default:100),update_every(default:10),inv_gamma(default:1.0),power(default:2 / 3),min_value(default:0.0),param_or_buffer_names_no_ema(default:set()),ignore_names(default:set()),ignore_startswith_names(default:set()),allow_different_devices(default:False),use_foreach(default:False),update_model_with_ema_every(default:None),update_model_with_ema_beta(default:0.0),forward_method_names(default:()),move_ema_to_online_device(default:False),coerce_dtype(default:False),lazy_init_ema(default:False).
- decay – EMA decay rate, forwarded to
on_load_checkpoint(trainer, pl_module, checkpoint)
Restore the EMA state from ema_model_state_dict if present.
on_save_checkpoint(trainer, pl_module, checkpoint)
Store the EMA state under ema_model_state_dict.
on_test_end(trainer, pl_module)
Put the online weights back after testing.
on_test_start(trainer, pl_module)
Swap the EMA weights in so testing runs on them.
on_train_batch_end(trainer, pl_module, outputs, batch, batch_idx)
Update the EMA weights once per true optimizer step.
on_train_start(trainer, pl_module)
Record the step counter training actually starts from.
on_validation_end(trainer, pl_module)
Put the online weights back after validation.
on_validation_start(trainer, pl_module)
Swap the EMA weights in so validation runs on them.
setup(trainer: Trainer, pl_module: LightningModule, stage: str)
Create the EMA copy of the model at the start of fit.
