adityachaubey commited on
Commit
9596f34
·
1 Parent(s): 629a11e

Add EMA class for exponential moving average weights management

Browse files
Files changed (1) hide show
  1. model.py +35 -1
model.py CHANGED
@@ -131,4 +131,38 @@ class Unet(torch.nn.Module):
131
  x7 = torch.cat([x2, x7], dim=1)
132
  x7 = self.resblock4(x7, t_emb)
133
 
134
- return self.outc(x7)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
131
  x7 = torch.cat([x2, x7], dim=1)
132
  x7 = self.resblock4(x7, t_emb)
133
 
134
+ return self.outc(x7)
135
+
136
+ #EMA weights
137
+ class EMA():
138
+ def __init__(self, model, decay=0.9999):
139
+ self.model = model
140
+ self.decay = decay
141
+ self.shadow = {}
142
+ self.backup = {}
143
+
144
+ def register(self):
145
+ for name , param in self.model.named_parameters():
146
+ if param.requires_grad:
147
+ self.shadow[name] = param.data.clone()
148
+
149
+ def upadate(self):
150
+ for name, param in self.model.named_parameters():
151
+ if param.requires_gard:
152
+ assert name in self.shadow
153
+ new_avg = self.decay*self.shadow[name] + ((1. - self.decay)*param)
154
+ self.shadow[name] = new_avg.clone()
155
+
156
+ def apply_shadow(self):
157
+ for name , param in self.model.named_parameters():
158
+ if param.requires_grad:
159
+ assert name in self.shadow
160
+ self.backup = param.data.clone()
161
+ param.data.copy_(self.shadow[name])
162
+
163
+ def restore(self):
164
+ for name, param in self.model.named_parameters():
165
+ if param.require_grad:
166
+ assert name in self.backup
167
+ param.data.copy_(self.backup[name])
168
+ self.backup = {}