MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / ema_update

Function ema_update

k_diffusion/utils.py:90–105  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

88
89@torch.no_grad()
90def 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
108class EMAWarmup:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected