Spaces:
Sleeping
Sleeping
Commit ·
9596f34
1
Parent(s): 629a11e
Add EMA class for exponential moving average weights management
Browse files
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 = {}
|