# ema-pytorch: a wrapper for tracking EMA weights in PyTorch training loops

> ema-pytorch wraps an existing torch.nn.Module and keeps a decayed copy of its parameters in sync, plus a post-hoc variant for synthesizing EMA models after training. It is a small library for people who already own their training loop.

**lucidrains/ema-pytorch** — A simple way to keep track of an Exponential Moving Average (EMA) version of your Pytorch model

- Repository: https://github.com/lucidrains/ema-pytorch
- Stars: 659 · Forks: 41
- Language: Python
- License: MIT
- Published: 2026-09-10 · Updated: 2026-09-10 · Language: en
- Canonical page: https://hysenlabs.com/projects/lucidrains-ema-pytorch

## The gap ema-pytorch fills between a raw model and a training framework

Exponential moving average of weights is a standard trick in diffusion and self-supervised training: you keep a second copy of the parameters that lags behind the trained ones, and you evaluate or sample from that copy instead. The recipe is short, but writing it correctly takes more care than it looks. You need a decay factor, a decision about when to start averaging, a decision about how often to run the copy, and a place to store the step counter so a resumed run does not restart the warmup.

ema-pytorch packages exactly that. It targets people who already have a training loop and do not want to adopt a framework to get one feature. The README describes it as "a simple way to keep track of an Exponential Moving Average (EMA) version of your pytorch model", and the wrapping style reflects that: you hand it a module, you call update() at the point in your loop where you would otherwise write the averaging math yourself. The package metadata lists a single runtime dependency, torch>=2.0, so it does not pull in a training framework, a logging stack or a configuration system.

That narrow scope is the point. If your loop is a plain for-loop over a DataLoader, this fits without restructuring anything.

## How the EMA wrapper tracks parameters, warmup and update cadence

The mechanism is a shadow copy. EMA(net, beta=...) holds an internal copy of your module, exposed as ema.ema_model, and update() moves that copy toward the live parameters by the decay factor. The README's example sets beta to 0.9999, which is the usual value for long diffusion runs; a smaller beta makes the copy track the live weights more closely, which is what the README's EMAModuleWrapper example does with beta=0.99 for a self-supervised setup.

Two arguments control when updates happen. update_after_step=100 means the wrapper ignores the first hundred calls to update(); this is the warmup logic the README attributes to @crowsonkb and says has been "validated for a number of projects now". update_every=10 means only every tenth call actually applies the averaging, which the README frames as a compute saving. The practical consequence is that the wrapper's internal step counter, not your epoch counter, decides the schedule. That is why the README recommends saving the entire wrapper rather than only the inner module: the counter lives in the wrapper, and a checkpoint that stores only ema.ema_model loses the warmup state.

The library also ships PostHocEMA, which follows Karras et al. (arXiv 2312.02696). Instead of one shadow model, you configure a tuple of sigma_rels such as (0.05, 0.28) and the wrapper writes checkpoints to checkpoint_folder at a fixed interval. After training, synthesize_ema_model(sigma_rel=0.15) produces an EMA model for a value you never trained with. The README requires at least two sigma_rels to synthesize a new one.

A third mode comes from the Switch EMA paper (arXiv 2402.09240): setting update_model_with_ema_every copies EMA weights back into the live model on a schedule, or you call ema.update_model_with_ema() manually at an epoch boundary. This is the only mode where the wrapper writes into your model rather than reading from it, and it changes training dynamics rather than just evaluation.

## Installing ema-pytorch and running a first update

The README gives one install line, with no extras, no optional dependency groups and no build step:

```bash
$ pip install ema-pytorch
```

After that, the smallest useful program wraps a module and calls update once per training step. The example below is the README's usage snippet: it wraps a Linear layer, mutates the live weights, calls update(), and then runs both the live model and the EMA model on the same input.

```python
import torch
from ema_pytorch import EMA

net = torch.nn.Linear(512, 512)

ema = EMA(
    net,
    beta = 0.9999,
    update_after_step = 100,
    update_every = 10,
)

with torch.no_grad():
    net.weight.copy_(torch.randn_like(net.weight))
    net.bias.copy_(torch.randn_like(net.bias))

ema.update()

data = torch.randn(1, 512)
output     = net(data)
ema_output = ema(data)
```

What you should see: ema(data) returns a tensor of the same shape as net(data), because the wrapper forwards to the EMA copy. With update_every=10 and a single update() call, the shadow weights will not have moved yet, since the counter has not reached the first update interval. If you want the EMA copy to move on every call while you are experimenting, pass update_every=1 and update_after_step=0.

To persist the state, the README's recommendation is to save the wrapper itself rather than the inner module, because the wrapper carries the number of steps taken:

```python
import torch

torch.save(ema.state_dict(), 'ema.pt')
```

Reload by constructing EMA around the same module shape and calling load_state_dict on the wrapper, not on ema.ema_model. If you only need the averaged weights for inference, ema.ema_model is the attribute to read.

## EMAModuleWrapper and target representation routing for nested models

The plain EMA class assumes one model and one shadow. EMAModuleWrapper handles a different shape of problem: self-supervised setups where a student branch should predict the output of a teacher branch, and the teacher is the EMA of the student. Rather than averaging the whole network, you name pairs of submodules, and the wrapper injects the EMA submodule's output into the online submodule's forward pass as a keyword argument.

The README's example maps 'branch_a.block1' to 'branch_b.block1' and 'branch_a.block2' to 'branch_b.block2'. The Block class in that example declares forward(self, x, ema_output=None) and returns a loss computed against that argument, so the wiring is explicit in the module signature. The default keyword name is ema_output; the README shows overriding it per mapping with an ema_kwarg key, and notes that you can also transform captured outputs or define custom mappings.

For multi-view training, where the student and teacher see different augmentations of the same sample, the README passes a separate input to the teacher side with ema_args or ema_kwargs, for example ema(student_input, ema_args=teacher_input).

This is the most opinionated part of the library. It assumes your submodules accept the injected keyword and that you are willing to write the loss inside the module that receives it. If your codebase passes a single tensor through a stack and computes losses at the top level, the plain EMA class is the better fit and EMAModuleWrapper is unnecessary machinery.

## Where ema-pytorch stops: checkpoints, distributed training and resumption

The README does not document distributed training. There is no mention of DistributedDataParallel, FSDP or any synchronization of the shadow copy across ranks, so if you train on multiple GPUs you have to decide yourself whether to wrap the model before or after DDP and whether each rank keeps its own EMA. The repository layout is a single package directory plus a tests folder, and the README covers the wrapper API rather than deployment concerns.

The post-hoc path has a cost the README states plainly: PostHocEMA writes checkpoints to checkpoint_folder every checkpoint_every_num_steps steps. That is disk I/O during training proportional to your model size divided by that interval, and it is the price of being able to synthesize an EMA for an arbitrary sigma_rel afterwards. If you only need one decay value, PostHocEMA is the wrong tool and the plain EMA class is cheaper.

The resumption story is the sharpest limitation. Because the warmup depends on the internal step count, restoring only ema.ema_model gives you averaged weights with no memory of how far the schedule had progressed. The README recommends saving the entire wrapper for this reason, but it does not document a resume recipe, so the loading path is on you. The same applies to optimizer state: the library does not touch it.

Finally, the package classifiers list Development Status :: 4 - Beta. The last push to the repository was on 2026-07-31, and the most recent release listed is 0.7.9 from 2025-12-19, while pyproject.toml declares version 0.8.3, so the published package version and the release list do not line up. Pin a version you have checked rather than tracking the default.

## ema-pytorch against timm's ModelEmaV2 and framework callbacks

The closest comparison is timm's ModelEmaV2, which solves the same problem with a different bias. timm's helper is a small class you instantiate alongside a model and call update on, and it is designed to be used with timm's own training scripts; the EMA copy is exposed as a module you evaluate directly. ema-pytorch differs in three ways the README makes visible. It adds warmup through update_after_step, so early noisy steps do not drag the average. It adds update_every for cadence control. And it adds the post-hoc synthesis path from Karras et al., which timm's helper does not have.

Framework callbacks are the other alternative. PyTorch Lightning users often reach for an EMA callback, which hooks into the training loop and needs no manual update() call. That is less code in your model file but more coupling: the averaging lives inside the framework's lifecycle, and moving to a plain loop means rewriting it. ema-pytorch makes the opposite trade. You call update() yourself, which is one more line per step, and in exchange the wrapper works in any loop, including ones with custom gradient accumulation or manual learning-rate schedules.

The honest summary is that ema-pytorch's edge is the post-hoc synthesis and the module-routing wrapper, not the basic averaging. If you only need a decayed copy of your weights and you already use timm or Lightning, the tool you have is probably sufficient.

## Conclusion

Adopt ema-pytorch if you write your own PyTorch training loop and want EMA weights without adding a training framework: the wrapper is a torch.nn.Module subclass, so ema(data) works anywhere net(data) did, and saving the whole wrapper preserves the step count the warmup logic depends on. Do not adopt it if you expect it to manage checkpoints, distributed replication or optimizer state; the README documents none of that, and it is not a Lightning callback. Before wiring it in, verify two things in your own code: that your update call site runs once per optimizer step rather than once per batch when you use gradient accumulation, and that your checkpoint loader reconstructs the EMA wrapper rather than only the raw model, since ema.ema_model alone does not carry the internal step counter.

## FAQ

### Is it better to use EMA or MA?

The README does not compare EMA with a simple moving average. It presents EMA as the method the library implements, with beta controlling the decay, and the post-hoc path from Karras et al. as an extension of it.

### What is EMA 12 and EMA 26?

Those terms do not appear in the README or the package metadata. In ema-pytorch the decay is set with the beta argument, for example beta=0.9999 in the README's main example.

### Is ChatGPT built on PyTorch?

The README and the package metadata say nothing about ChatGPT. What they do state is that ema-pytorch requires torch>=2.0 and wraps a torch.nn.Module.

### What is EMA?

In this library, EMA means an exponential moving average copy of your PyTorch model, described in the README as "a simple way to keep track of an Exponential Moving Average (EMA) version of your pytorch model". The copy lives at ema.ema_model and is updated by calling ema.update().

## Sources

- [Issues](https://github.com/lucidrains/ema-pytorch/issues)
- [License: MIT](https://github.com/lucidrains/ema-pytorch/blob/main/LICENSE)
- [lucidrains/ema-pytorch on GitHub](https://github.com/lucidrains/ema-pytorch)
- [README](https://github.com/lucidrains/ema-pytorch/blob/main/README.md)
- [Releases](https://github.com/lucidrains/ema-pytorch/releases)

---

Hysen Labs editorial analysis, written from the project's own repository and release notes. Cite the canonical page: https://hysenlabs.com/projects/lucidrains-ema-pytorch
