espnet3.components.callbacks.vendored_ema.EMA
espnet3.components.callbacks.vendored_ema.EMA
class espnet3.components.callbacks.vendored_ema.EMA(model: Module, ema_model: Module | Callable[[], Module] | None = None, beta=0.9999, update_after_step=100, update_every=10, inv_gamma=1.0, power=0.6666666666666666, min_value=0.0, param_or_buffer_names_no_ema: set[str] = {}, ignore_names: set[str] = {}, ignore_startswith_names: set[str] = {}, include_online_model=True, allow_different_devices=False, use_foreach=False, update_model_with_ema_every=None, update_model_with_ema_beta=0.0, forward_method_names: tuple[str, ...] = (), move_ema_to_online_device=False, coerce_dtype=False, lazy_init_ema=False)
Bases: Module
Implements exponential moving average shadowing for your model.
Utilizes an inverse decay schedule to manage longer term training runs. By adjusting the power, you can control how fast EMA will ramp up to your specified beta.
@crowsonkb’s notes on EMA Warmup:
If gamma=1 and power=1, implements a simple average. gamma=1, power=2/3 are good values for models you plan to train for a million or more steps (reaches decay factor 0.999 at 31.6K steps, 0.9999 at 1M steps), gamma=1, power=3/4 for models you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999 at 215.4k steps).
- Parameters:
- inv_gamma (float) – Inverse multiplicative factor of EMA warmup. Default: 1.
- power (float) – Exponential factor of EMA warmup. Default: 2/3.
- min_value (float) – The minimum EMA decay rate. Default: 0.
Initialize EMA state from vendored ema-pytorch code.
This implementation is copied from lucidrains/ema-pytorch. Refer to that project’s documentation for the complete EMA API and behavioral details; ESPnet keeps this copy to avoid a runtime dependency.
add_to_optimizer_post_step_hook(optimizer)
copy_params_from_ema_to_model()
copy_params_from_model_to_ema()
eval()
Set the module in evaluation mode.
This has an effect only on certain modules. See the documentation of particular modules for details of their behaviors in training/evaluation mode, i.e. whether they are affected, e.g. Dropout, BatchNorm, etc.
This is equivalent with self.train(False).
See locally-disable-grad-doc for a comparison between .eval() and several similar mechanisms that may be confused with it.
- Returns: self
- Return type: Module
forward_eval(*args, **kwargs)
get_buffers_iter(model)
get_current_decay()
get_params_iter(model)
init_ema(ema_model: Module | None = None)
property model
restore_ema_model_device()
update()
update_model_with_ema(decay=None)
update_moving_average(ma_model, current_model, current_decay=None)
