EMA¶
Exponential moving average of model weights, updated once per optimizer step with a warmup ramp. Validation and checkpointing run under the averaged weights — these are the weights you evaluate and ship.
ema
¶
Exponential Moving Average of Model Weights¶
Flow-matching and diffusion training produce noisy gradient estimates: the loss regresses a velocity target that is only conditionally determined by the network input, so consecutive iterates bounce around the optimum rather than settling into it. Averaging weights over training suppresses that jitter and is standard practice for this model family — the averaged weights are what you evaluate and ship, while the raw weights keep training.
Usage:
ema = EMA(model, decay=0.999)
for batch in loader:
loss = objective.compute_loss(batch)
loss.backward()
optimizer.step()
ema.update() # after every optimizer step
with ema.average_parameters():
metrics = evaluate(model) # model temporarily holds EMA weights
EMA
¶
Maintains an exponential moving average of a model's parameters.
The shadow weights follow
with decay \(d\). Early in training the shadow is dominated by its
(arbitrary) initial value, so warmup ramps the effective decay up from
0 following \(\min(d, (1 + n) / (10 + n))\) at step \(n\). This
is the schedule used by most diffusion implementations and it removes the
need to tune a separate "start EMA at step N" threshold.
Buffers (e.g. BatchNorm running statistics) are copied rather than averaged, matching the reference implementations.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model
|
Module
|
Model whose parameters are tracked. |
required |
decay
|
float
|
Target decay rate. Typical values are 0.999 (short runs) to 0.9999 (long runs). Higher means smoother and slower to adapt. |
0.999
|
warmup
|
bool
|
Ramp the effective decay early in training (default: True). |
True
|
device
|
Optional[device]
|
Optional device for the shadow copy. Defaults to keeping each shadow tensor on the same device as its source parameter. |
None
|
Example
model = nn.Linear(2, 2) ema = EMA(model, decay=0.9, warmup=False) ema.update() with ema.average_parameters(): ... _ = model(torch.zeros(1, 2)) # uses averaged weights
Source code in flowpde/trainers/ema.py
37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 | |
current_decay
property
¶
Effective decay at the current step, accounting for warmup.
update(model=None)
¶
Update the shadow weights. Call once after every optimizer step.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model
|
Optional[Module]
|
Optionally update from a different model instance; defaults to the model this EMA was constructed with. |
None
|
Source code in flowpde/trainers/ema.py
copy_to(model=None)
¶
Write the averaged weights into model, permanently.
Source code in flowpde/trainers/ema.py
store(model=None)
¶
Stash the model's live weights so they can be restored later.
Source code in flowpde/trainers/ema.py
restore(model=None)
¶
Put back the weights saved by store().
Source code in flowpde/trainers/ema.py
average_parameters(model=None)
¶
Temporarily swap the averaged weights into the model.
The live training weights are restored on exit, including when the body raises, so a failed validation pass cannot corrupt training.
Source code in flowpde/trainers/ema.py
to(device)
¶
state_dict()
¶
Serializable state for checkpointing.
Source code in flowpde/trainers/ema.py
load_state_dict(state)
¶
Restore state saved by state_dict().