# ema

Python API: `luxonis_train.callbacks.ema`

Keeps an exponential moving average of the weights.

Validation during a fit runs with the average weights, so its visualizations also come from these weights. After the fit, the
model keeps the average weights. A checkpoint that is not weights-only holds the average weights as the model weights. An export
from such a checkpoint thus uses them too.

## Classes

### EMACallback

Callback that keeps an exponential moving average of the weights.

The callback keeps the average in a
[ModelEma](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
It changes the model at these points of a run:

 * [on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
   creates the average from the current weights.
 * [on_train_batch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
   moves the average toward the trained weights once in each gradient accumulation window.
 * Validation and test run with the average weights, when the average exists. The callback keeps a copy of the model weights when
   the loop starts, and loads the copy back when the loop ends.
 * When training ends, the model keeps the average weights.
 * Each checkpoint that is not weights-only holds the average weights as the model weights. The state of the callback in the
   checkpoint also holds the average and the update count, so a resumed fit continues the average.

While
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model, the callback does not swap any weights.

[LuxonisLightningModule.configure_callbacks](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
builds a new instance of this callback for each call of the trainer, such as `trainer.fit` or `trainer.test`. Thus a
`trainer.test` call after the fit has no average, and the test runs with the current weights of the model.

The config moves this callback to the front of `trainer.callbacks`, so that it runs before the other callbacks of the config.

#### Methods

##### init

```python
def __init__(decay: float = 0.5, use_dynamic_decay: bool = True, decay_tau: float = 2000):
```

Initialize the callback with the options of the average.

The callback creates the average in
[on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md),
not here.

Parameters

 * `decay` (`float`): The largest decay of the average. A value near `1` moves the average slowly. The default `0.5` is far lower
   than the `0.9999` default of
   [ModelEma](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
 * `use_dynamic_decay` (`bool`): When `True`, the decay grows from `0` toward `decay` as the updates add up. See
   [ModelEma.update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
 * `decay_tau` (`float`): The time constant of the dynamic decay, in updates.

##### load_state_dict

```python
def load_state_dict(state_dict: dict[str, Any]):
```

Store a checkpoint average for the next
[on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).

Lightning calls this method when it restores a checkpoint, with the dictionary that
[state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
returned. The method reads the value under `"ema_state_dict"`, or under `"state_dict"` when that key is missing. It keeps that
value when the value is a mapping. It also keeps the value under `"updates"` when that value is an `int`. It ignores a value of
another type and an empty `state_dict`. The current average does not change.

> **Example**
> Before a fit, the callback has no average, so its state stays empty:

```pycon
>>> import torch
>>> callback = EMACallback()
>>> callback.load_state_dict(
...     {"ema_state_dict": {"weight": torch.ones(1)}, "updates": 7}
... )
>>> callback.state_dict()
{}
```

Parameters

 * `state_dict` (`dict[str, Any]`): The state of the callback.

##### on_fit_start

```python
def on_fit_start(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Create the average from the current weights of the model.

Lightning calls this hook at the start of `trainer.fit`. The hook builds a new
[ModelEma](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
from `pl_module` with the options of the callback. When
[load_state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
or
[on_load_checkpoint](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
stored an average before, the hook copies it into the new average:

 * It ignores the entries that losses, metrics, and visualizers keep for their node, in both averages.
 * It logs a warning for the keys that the stored average misses, the keys it has in excess, and the keys with another shape.
 * It keeps the new values for the missing keys and for the keys with another shape.
 * It moves the stored tensors to the device of the new average.
 * It sets the update count when a count was stored.

It then clears the stored average and count.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The model to average.
   [ModelEma](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
   leaves it in training mode.

##### on_load_checkpoint

```python
def on_load_checkpoint(trainer: pl.Trainer, pl_module: pl.LightningModule, callback_state: dict):
```

Store the checkpoint weights for the next
[on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).

Lightning calls this hook when it restores a checkpoint that has a `callbacks` key, before it calls
[load_state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
Lightning passes the whole checkpoint, not the state of the callback. The hook reads it like
[load_state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
The checkpoint has no `"ema_state_dict"` key at the top level, so the hook stores its `"state_dict"`: the average weights, when
this callback saved the checkpoint.
[load_state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
then replaces that value with the `"ema_state_dict"` of the callback state, when the checkpoint has one.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The model. Unused.
 * `callback_state` (`dict`): The whole checkpoint, despite the name.

##### on_save_checkpoint

```python
def on_save_checkpoint(trainer: pl.Trainer, pl_module: pl.LightningModule, checkpoint: dict):
```

Store the average as the model weights of the checkpoint.

Lightning calls this hook when it saves a checkpoint that is not weights-only, after it collects the output of
[state_dict](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
When the average exists, the hook replaces `checkpoint["state_dict"]` with the average. A model that loads the checkpoint thus
gets the average weights.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The model. Unused.
 * `checkpoint` (`dict`): The checkpoint that Lightning saves. The hook changes it in place.

##### on_test_end

```python
def on_test_end(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Load the kept weights back into the model after the test.

Lightning calls this hook when the test loop ends. The hook loads the kept weights into `pl_module`, like
[on_validation_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
It does nothing when no copy exists, or while
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The tested model.

##### on_test_epoch_start

```python
def on_test_epoch_start(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Load the average weights into the model for the test.

Lightning calls this hook at the start of each test epoch. The hook swaps the weights like
[on_validation_epoch_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md):
it keeps a deep copy of the weights of `pl_module`, then loads the average when it exists. While
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model, the hook does nothing.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The model to test.

##### on_train_batch_end

```python
def on_train_batch_end(trainer: pl.Trainer, pl_module: pl.LightningModule, outputs: STEP_OUTPUT, batch: Any, batch_idx: int):
```

Move the average one step toward the trained weights.

Lightning calls this hook after each training batch. The hook calls
[ModelEma.update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
when the average exists and `batch_idx` is a multiple of `trainer.accumulate_grad_batches`. Thus the average moves once in each
gradient accumulation window, on the first batch of the window.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. The hook reads its `accumulate_grad_batches`.
 * `pl_module` (`pl.LightningModule`): The model whose weights the average moves toward.
 * `outputs` (`STEP_OUTPUT`): The output of the training step. Unused.
 * `batch` (`Any`): The batch. Unused.
 * `batch_idx` (`int`): The index of the batch in the epoch.

##### on_train_end

```python
def on_train_end(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Load the average weights into the model when training ends.

Lightning calls this hook once when `trainer.fit` ends. The hook swaps the weights like
[on_validation_epoch_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md),
but it does not load the trained weights back. Thus the model keeps the average weights after the fit. While
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model, the hook does nothing.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The trained model.

##### on_validation_end

```python
def on_validation_end(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Load the kept weights back into the model after validation.

Lightning calls this hook when the validation loop ends. The hook loads the kept weights into `pl_module`. During a fit, these are
the trained weights. The hook does nothing when no copy exists, or while
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model. The copy stays in the callback.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The validated model.

##### on_validation_epoch_start

```python
def on_validation_epoch_start(trainer: pl.Trainer, pl_module: pl.LightningModule):
```

Load the average weights into the model for validation.

Lightning calls this hook at the start of each validation epoch. The hook keeps a deep copy of the state dictionary of
`pl_module`. It then loads the average into `pl_module` when the average exists. While
[replace_weights](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/utils/export_utils.md)
holds explicit weights in the model, the hook does nothing.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`pl.LightningModule`): The model to validate.

##### state_dict

```python
def state_dict(self) -> dict[str, Any]:
```

Return the state of the callback for a checkpoint.

Lightning calls this method when it saves a checkpoint that is not weights-only. It stores a non-empty result in the `callbacks`
entry of the checkpoint, under the state key of the callback. It does not store an empty result.

Returns

 * `dict[str, Any]`: An empty dictionary when the average does not exist yet. Otherwise a dictionary with two keys. *
   `"ema_state_dict"` holds the average, without the entries that losses, metrics, and visualizers keep for their node.
    * `"updates"` holds the number of updates of the average.

#### Attributes

##### ema

The
[ModelEma](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
that holds the average.

Raises

 * `ValueError`: When
   [on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
   has not created the average yet.

### ModelEma

Exponential moving average of the state dictionary of a model.

The average covers every entry of `model.state_dict()`: the parameters and the persistent buffers. It lives in a plain dictionary,
not in registered buffers, so `ModelEma.state_dict()` does not return it.
[EMACallback](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
creates and updates it.

#### Methods

##### init

```python
def __init__(model: pl.LightningModule, decay: float = 0.9999, use_dynamic_decay: bool = True, decay_tau: float = 2000):
```

Copy the state dictionary of `model` as the first average.

The method calls `model.eval()`, makes a deep copy of `model.state_dict()`, and then calls `model.train()`. Thus `model` is in
training mode afterwards, whatever its mode was before. The copy does not require gradients.

Parameters

 * `model` (`pl.LightningModule`): The model to average.
 * `decay` (`float`): The largest decay d. A value near `1` moves the average slowly.
 * `use_dynamic_decay` (`bool`): When `True`, the decay grows from `0` toward `decay` as the updates add up. See
   [update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md).
 * `decay_tau` (`float`): The time constant τ of the dynamic decay, in updates. A larger value makes the decay grow more slowly.

##### update

```python
def update(model: pl.LightningModule):
```

Move the average one step toward the state of `model`.

The method adds `1` to `updates` and computes the decay dt of this step. With `use_dynamic_decay`, the decay grows from `0` toward
`decay` as the updates add up, so the first averages follow the model closely:

dt = d(1 − e − t ⁄ τ)

Here t is `updates`, d is `decay`, and τ is `decay_tau`. Without `use_dynamic_decay`, dt = d. The method then changes each
floating point entry θema of the average in place:

θema ← dt θema + (1 − dt) θ

Here θ is the entry of `model.state_dict()` with the same key. The method copies each entry of another type, such as the
`num_batches_tracked` counter of a batch norm. It skips a key that `model` does not have. No gradients flow through the update.

> **References**
> * Source: adapted from [timm model_ema.py](https://github.com/huggingface/pytorch-image-models/blob/main/timm/utils/model_ema.py) ([Apache License 2.0](https://github.com/huggingface/pytorch-image-models/tree/main?tab=Apache-2.0-1-ov-file#readme)).

> **Examples**
> A fixed decay of `0.5` moves the average halfway:

```pycon
>>> import lightning.pytorch as pl
>>> import torch
>>> model = pl.LightningModule()
>>> model.weight = torch.nn.Parameter(torch.zeros(2))
>>> ema = ModelEma(model, decay=0.5, use_dynamic_decay=False)
>>> _ = model.weight.data.fill_(1.0)
>>> ema.update(model)
>>> ema.state_dict_ema["weight"].tolist()
[0.5, 0.5]
```

The dynamic decay is smaller for the first updates, so the average moves farther toward the model:

```pycon
>>> _ = model.weight.data.zero_()
>>> ema = ModelEma(model, decay=0.5, decay_tau=1.0)
>>> _ = model.weight.data.fill_(1.0)
>>> ema.update(model)
>>> round(ema.state_dict_ema["weight"][0].item(), 3)
0.684
>>> ema.updates
1
```

Parameters

 * `model` (`pl.LightningModule`): The model whose current state the average moves toward.

#### Attributes

##### state_dict_ema

The average, with the keys of `model.state_dict()`.

##### updates

The number of updates of the average. Each
[update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
call adds `1`.
[EMACallback.on_fit_start](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/ema.md)
sets it from a checkpoint that holds a count.
