MCPcopy Create free account
hub / github.com/dome272/Diffusion-Models-pytorch / EMA

Class EMA

modules.py:7–32  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class EMA:
8 def __init__(self, beta):
9 super().__init__()
10 self.beta = beta
11 self.step = 0
12
13 def update_model_average(self, ma_model, current_model):
14 for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()):
15 old_weight, up_weight = ma_params.data, current_params.data
16 ma_params.data = self.update_average(old_weight, up_weight)
17
18 def update_average(self, old, new):
19 if old is None:
20 return new
21 return old * self.beta + (1 - self.beta) * new
22
23 def step_ema(self, ema_model, model, step_start_ema=2000):
24 if self.step < step_start_ema:
25 self.reset_parameters(ema_model, model)
26 self.step += 1
27 return
28 self.update_model_average(ema_model, model)
29 self.step += 1
30
31 def reset_parameters(self, ema_model, model):
32 ema_model.load_state_dict(model.state_dict())
33
34
35class SelfAttention(nn.Module):

Callers 1

trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected