Incorporates updated model parameters into an exponential moving averaged version of a model. It should be called after each optimizer step.
(model, averaged_model, decay)
| 88 | |
| 89 | @torch.no_grad() |
| 90 | def ema_update(model, averaged_model, decay): |
| 91 | """Incorporates updated model parameters into an exponential moving averaged |
| 92 | version of a model. It should be called after each optimizer step.""" |
| 93 | model_params = dict(model.named_parameters()) |
| 94 | averaged_params = dict(averaged_model.named_parameters()) |
| 95 | assert model_params.keys() == averaged_params.keys() |
| 96 | |
| 97 | for name, param in model_params.items(): |
| 98 | averaged_params[name].lerp_(param, 1 - decay) |
| 99 | |
| 100 | model_buffers = dict(model.named_buffers()) |
| 101 | averaged_buffers = dict(averaged_model.named_buffers()) |
| 102 | assert model_buffers.keys() == averaged_buffers.keys() |
| 103 | |
| 104 | for name, buf in model_buffers.items(): |
| 105 | averaged_buffers[name].copy_(buf) |
| 106 | |
| 107 | |
| 108 | class EMAWarmup: |
nothing calls this directly
no outgoing calls
no test coverage detected