Model Averaging — EMA, SWA¶
Overview¶
Averaging models over training is a free ensemble: it smooths the loss landscape, de-noises the last-epoch chatter, and routinely adds 0.5-2% accuracy or BLEU/ROUGE for nothing at inference time.
- EMA (Exponential Moving Average):
ema = decay * ema + (1-decay) * current, updated every step. Best for final deployment weights. - SWA (Stochastic Weight Averaging): uniform mean of checkpoints over the last few epochs. Better for training-end quality and flat minima.
- EMA decay example: 0.999-0.9999 for long jobs; 0.99 for short bursts.
💡 Rule: use the EMA weights at eval/inference; keep the live weights for training (they adapt faster). Never train the EMA copy itself.
EMA — the Standard Implementation¶
import copy, torch, torch.nn as nn
class EMA:
def __init__(self, model, decay=0.999):
self.decay = decay
self.ema = copy.deepcopy(model) # shadow copy
for p in self.ema.parameters():
p.requires_grad_(False)
@torch.no_grad()
def update(self, model):
ema_p = dict(self.ema.named_parameters())
for name, p in model.named_parameters():
if name in ema_p:
ema_p[name].mul_(self.decay).add_(p.detach(), alpha=1 - self.decay)
def state_dict(self):
return self.ema.state_dict()
ema = EMA(model, decay=0.999)
for step in range(100):
loss = (model(torch.randn(8, 4))).abs().mean()
loss.backward(); opt.step()
ema.update(model) # shadow follows live model
# switch to EMA at eval:
ema_model = ema.ema.eval()
with torch.no_grad():
print(ema_model(torch.randn(1, 4)).shape)
⚠️
deepcopyneeds the model on CPU or the same device — either deepcopy before.cuda()or move after copy.
Buffers & EMA (Batchnorm gotcha)¶
Copying only parameters leaves BatchNorm running stats stale. Copy buffers too:
@torch.no_grad()
def ema_update(self, model):
for name, p in model.named_parameters():
if name in self.ema_names:
self.ema_params[name].mul_(self.decay).add_(p.detach(), alpha=1 - self.decay)
for name, b in model.named_buffers():
if name in self.ema_buffers:
self.ema_buffers[name].copy_(b) # buffers follow live copy
💡 Many opt for: after training, re-run a few eval forwards on EMA model so BatchNorm recalibrates (
model.train(False)+ forward pass on real data).
SWA — Stochastic Weight Averaging¶
from torch.optim.swa_utils import AveragedModel, SWALR
swa_model = AveragedModel(model) # averages params over time
swa_sched = SWALR(opt, swa_lr=0.05) # low, constant LR for the SWA phase
for epoch in range(10):
train_one_epoch()
if epoch >= 7: # start averaging near the end
swa_model.update_parameters(model)
swa_sched.step()
torch.optim.swa_utils.update_bn(swa_model, val_loader) # recalibrate BN stats
EMA vs SWA — When to Use What¶
| EMA | SWA | |
|---|---|---|
| Update frequency | every step | per epoch/checkpoint |
| Effective ensemble | infinite memory | last-N checkpoints only |
| Final weights | long-run smoothed | flattened minima |
| Best for | chat/LLM deployment, RL value nets | vision/classical-generalization tasks |
| Decay/lr | 0.99-0.9999 | SWA phase LR ~ peak/10 |
✅ EMA is the modern default for LLM/RL deployment; SWA shines on classification deadlines with fewer training steps.
Pitfalls¶
- Use EMA at inference, not in training metrics — train metrics reflect fast weights.
torch.compile+ EMA: compile the EMA copy separately or keep EMA un-compiled (shadows don't benefit).- Don't EMA-then-distill: EMA-of-EMA compounds; one EMA layer is enough.
- Checkpoint both: save
modelandema.state_dict()— resume keeps both consistent.
ckpt = {"model": model.state_dict(), "ema": ema.state_dict(), "step": step}
torch.save(ckpt, "ckpt.pt")
Key Takeaways¶
- EMA/SWA are free ensembles: no extra forward passes, small memory (one shadow copy).
- EMA: per-step decay, best for deployment; SWA: end-phase averaging, best for final checkpoints.
- Copy buffers (BN stats) in EMA; recalibrate BN after SWA.
- Serve EMA weights at inference; keep the fast live weights during training.
- Save both in checkpoints for clean resumability.