Checkpoints, Export and Inference
Exponential moving average of weights
Keep a second, slowly-updated copy of your weights that trails the training ones — it is smoother, it usually evaluates better, and it costs one cheap blend per step.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Keep a second copy of the weights that follows the training ones slowly, and ship that copy instead.
Think of a shop's daily sales. Today's figure jumps around — a festival, a rainy afternoon, a wedding order. Nobody makes decisions on today's number. They watch a rolling average, which moves smoothly and tells you where business actually is.
Training weights are today's number. An exponential moving average, EMA for short, is the rolling average of them.
Why it helps
Near the end of training the weights do not settle. They circle around a good region, bouncing because each batch pulls in a slightly different direction.
Stopping at any single moment means stopping at one point on that bouncing path. The average of the path sits nearer the centre of the good region than most individual points on it.
So the averaged copy is usually a little better than the live one, for free, with no change to how you train.
How it works
every training step:
live weights ---- one normal training step ----> live weights (jumpy)
|
v blend a tiny bit of the live weights into the average
ema weights = 0.999 x ema + 0.001 x live (smooth)
at the end: evaluate and ship the EMA copyThe blend fraction is called the decay. A decay of 0.999 means the average keeps 99.9% of itself and takes 0.1% of the new weights. Higher decay means smoother and slower to react.
A real example you have seen
A cricketer's batting average. One duck does not sink it, one century does not save it. It moves slowly, and it describes the player far better than the last innings does.
Remember this
- EMA is a second copy of the weights, updated by a small blend each step.
- It is smoother than the live weights and usually scores a little better.
- Costs one extra copy of the model in memory, and a cheap blend per step.
What to learn next
- TorchScript: tracing vs scripting — packaging the weights you decided to ship.
- Early stopping and keeping the best model — the other way to pick which weights to keep.
- BatchNorm in PyTorch — why the running statistics need a second pass.
Developer — Code and libraries.
Setup
pip install torchRuns on CPU in a second. PyTorch ships EMA inside torch.optim.swa_utils, so there is nothing to write yourself.
The whole thing, with a measurable result
import torch
import torch.nn as nn
from torch.optim.swa_utils import AveragedModel, get_ema_multi_avg_fn
torch.manual_seed(0)
X = torch.randn(512, 6)
w_true = torch.randn(6, 1)
y = X @ w_true + 0.3 * torch.randn(512, 1) # noisy labels
Xtr, ytr, Xva, yva = X[:400], y[:400], X[400:], y[400:]
model = nn.Linear(6, 1)
ema = AveragedModel(model, multi_avg_fn=get_ema_multi_avg_fn(0.95)) # 0.95 = decay
opt = torch.optim.SGD(model.parameters(), lr=0.35) # deliberately jumpy
def val(m):
with torch.no_grad():
return nn.functional.mse_loss(m(Xva), yva).item()
for step in range(120):
idx = torch.randint(0, 400, (8,)) # tiny batches -> noisy steps
opt.zero_grad(set_to_none=True)
nn.functional.mse_loss(model(Xtr[idx]), ytr[idx]).backward()
opt.step()
ema.update_parameters(model) # one cheap blend per step
if step % 30 == 29:
print(f"step {step+1:3d} live weights {val(model):.4f} ema weights {val(ema):.4f}")step 30 live weights 0.1390 ema weights 0.2099 step 60 live weights 0.1705 ema weights 0.0761 step 90 live weights 0.1695 ema weights 0.0741 step 120 live weights 0.1170 ema weights 0.0711
Read those four rows carefully, because they contain both halves of the story.
At step 30 the EMA is worse — 0.2099 against 0.1390. It is still carrying the weights from early training, which were bad. The average lags.
From step 60 onward it is decisively better, and it stays better while the live model bounces between 0.117 and 0.171. At step 120 the EMA validation loss is under two-thirds of the live one.
That early lag is why EMA is a late-training tool. Evaluate it before it has caught up and you will conclude it does not work.
Note the decay here is 0.95, not the 0.999 you see in real projects. With 120 steps a decay of 0.999 would barely move off its starting point. The rule of thumb is that an EMA has a memory of roughly 1 / (1 - decay) steps: 20 steps at 0.95, and 1000 steps at 0.999. Pick decay against your run length, not by copying a number.
Saving and shipping the average
import torch
import torch.nn as nn
from torch.optim.swa_utils import AveragedModel, get_ema_multi_avg_fn
torch.manual_seed(0)
model = nn.Linear(4, 2)
ema = AveragedModel(model, multi_avg_fn=get_ema_multi_avg_fn(0.9))
for _ in range(5):
with torch.no_grad():
model.weight.add_(torch.randn_like(model.weight) * 0.1)
ema.update_parameters(model)
print("AveragedModel keys:", list(ema.state_dict())[:4])
print("the inner model's keys:", list(ema.module.state_dict()))
torch.save({"model": model.state_dict(),
"ema": ema.module.state_dict()}, "both.pt") # .module, not the wrapper
plain = nn.Linear(4, 2) # deploy target: no wrapper
plain.load_state_dict(torch.load("both.pt", weights_only=True)["ema"])
print("loaded the EMA weights into a plain model:", plain.weight.shape)AveragedModel keys: ['n_averaged', 'module.weight', 'module.bias'] the inner model's keys: ['weight', 'bias'] loaded the EMA weights into a plain model: torch.Size([2, 4])
AveragedModel is a wrapper, so its own state_dict() prefixes everything with module. — the same prefix problem as DDP checkpoints. Save ema.module.state_dict() and the file drops straight into a plain model.
Save both sets of weights. You need the live ones to resume training and the EMA ones to serve. They are different models from here on.
The BatchNorm caveat
Averaging touches parameters. It does not fix the running statistics that BatchNorm collected, which belong to the live weights and can be a poor match for the averaged ones.
For models with BatchNorm, run one pass over the training data to recompute them:
from torch.optim.swa_utils import update_bn
update_bn(train_loader, ema) # one forward-only pass, no gradientsSkip this and your EMA model can evaluate far worse than the live one, which sends people off hunting an imaginary bug in the averaging.
Where you have already met this
EMA is not a niche trick.
- Diffusion models. Nearly every image generator publishes EMA weights. Sampling from the live weights gives visibly worse images.
- Self-supervised learning. BYOL and MoCo use an EMA copy as the "teacher" that the live network learns to match.
- Semi-supervised learning. The Mean Teacher method is exactly this idea, with the average producing the training targets.
- Detection and segmentation. YOLO and its relatives ship EMA weights by default.
Common mistakes
Evaluating the EMA too early. The step-30 row above. Give it several times 1 / (1 - decay) steps before you judge it.
Copying decay=0.9999 into a 500-step fine-tune. That average has a memory of 10,000 steps. It will still be close to your initialisation when the run ends.
Saving ema.state_dict() rather than ema.module.state_dict(). Every key gains module. and the file will not load into your deployment model.
Forgetting update_bn on a BatchNorm model. The most common reason EMA "does not work".
Calling update_parameters before optimizer.step(). You would then average the weights from before the update, lagging by one step forever. Order matters: step, then update.
Expecting a big win. EMA typically buys a small, reliable improvement. It is not a substitute for a better learning-rate schedule.
Try it yourself
Change the decay in ema_demo.py to 0.5, then 0.99, and print every 10 steps. Watch the crossover point — where the EMA overtakes the live model — move later as decay rises.
What to learn next
- TorchScript: tracing vs scripting — packaging the weights you decided to ship.
- Early stopping and keeping the best model — the other way to pick which weights to keep.
- BatchNorm in PyTorch — why the running statistics need a second pass.
Researcher — Mathematics and papers.
The update, and what it estimates
For parameters $\theta_t$ at step $t$ and decay $\rho$, the EMA is
$$ \theta^{\text{ema}}t = \rho\,\theta^{\text{ema}}{t-1} + (1-\rho)\,\theta_t $$
where $\rho \in (0,1)$. Unrolling gives $\theta^{\text{ema}}t = (1-\rho)\sum{k=0}^{t-1}\rho^k \theta_{t-k}$ plus a decaying term from the initialisation. The weights $\rho^k$ sum to 1 in the limit, and the effective averaging window — the centre of mass of that geometric kernel — is $\rho/(1-\rho)$, which is the $1/(1-\rho)$ rule of thumb above. The initialisation term decays as $\rho^t$, so the estimate carries a bias of order $\rho^t$; this is the measured lag at step 30 in the output, and it is the same bias that Adam's bias correction removes in its moment estimates. Implementations that care about early behaviour use a warmup decay such as $\min(\rho, (1+t)/(10+t))$, which is what the diffusion literature adopted.
Why the average is better than its members: near convergence, SGD with constant learning rate does not converge to a point but to a stationary distribution around a minimum, with covariance proportional to $\eta \Sigma / B$ for learning rate $\eta$, gradient covariance $\Sigma$ and batch size $B$. Averaging reduces the variance of that iterate without moving the mean much, provided the loss is approximately quadratic over the region the iterates explore. That premise fails when the trajectory is still travelling — early training, or immediately after a learning-rate drop — which bounds when EMA helps.
Relation to SWA and to Polyak averaging
Polyak–Ruppert averaging takes the uniform mean $\frac{1}{T}\sum_t \theta_t$ and, for convex problems with decaying step sizes, achieves the optimal asymptotic rate — the iterate average converges faster than the iterates. EMA is the exponentially-weighted variant, which forgets early iterates and is therefore usable during a run rather than only at its end. SWA (Izmailov et al., 2018) is the uniform average taken over a cyclic or high-constant learning-rate schedule, deliberately sampling widely separated points in the basin; PyTorch implements both through the same AveragedModel, differing only in avg_fn. SWA's reported gains come with the same BatchNorm caveat and the same update_bn remedy.
The teacher-student use is a different mechanism with the same update. In BYOL, MoCo and Mean Teacher, the EMA copy is not only evaluated — its outputs form the training target, and the slow update is what prevents representational collapse by making the target a slowly-moving objective rather than a co-adapting one.
Cost
Memory: one extra copy of the parameters, $4P$ bytes in float32, which is small against Adam's $8P$ of moments. Compute: one fused multiply-add per parameter per step, negligible against a backward pass — PyTorch uses torch._foreach_lerp_ so the whole model is a handful of kernel launches. Under FSDP the EMA copy shards alongside the parameters, so the cost divides by the world size too. The real operational cost is remembering to checkpoint, evaluate and deploy the second set of weights.
References
- Polyak and Juditsky (1992), Acceleration of Stochastic Approximation by Averaging, SIAM J. Control Optim. — the original averaging result.
- Izmailov et al. (2018), Averaging Weights Leads to Wider Optima and Better Generalization — SWA, and the wide-minima argument.
- Tarvainen and Valpola (2017), Mean teachers are better role models — EMA weights as a training target.
- Grill et al. (2020), Bootstrap Your Own Latent (BYOL) — the EMA target network in self-supervised learning.
- Ho et al. (2020), Denoising Diffusion Probabilistic Models — EMA as standard practice in generative image models.
What to learn next
- TorchScript: tracing vs scripting — packaging the weights you decided to ship.
- Early stopping and keeping the best model — the other way to pick which weights to keep.
- BatchNorm in PyTorch — why the running statistics need a second pass.