Spaces:
Sleeping
Sleeping
Commit ·
9122c78
1
Parent(s): 9596f34
Add DDIM sampling function for image generation
Browse files
model.py
CHANGED
|
@@ -165,4 +165,27 @@ class EMA():
|
|
| 165 |
if param.require_grad:
|
| 166 |
assert name in self.backup
|
| 167 |
param.data.copy_(self.backup[name])
|
| 168 |
-
self.backup = {}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 165 |
if param.require_grad:
|
| 166 |
assert name in self.backup
|
| 167 |
param.data.copy_(self.backup[name])
|
| 168 |
+
self.backup = {}
|
| 169 |
+
|
| 170 |
+
#DDIM
|
| 171 |
+
@torch.no_grad()
|
| 172 |
+
def sample_ddim(model, T,img_size, batch_size=16, channels=3,ddim_step=50, device='cpu'):
|
| 173 |
+
model.eval()
|
| 174 |
+
alphas_cumprod_device = alphas_cumprod.to(device)
|
| 175 |
+
timestep = torch.linspace(0, T-1, steps=ddim_step, dtype=torch.long, device=device)
|
| 176 |
+
img = torch.randn(batch_size, channels, img_size, img_size).to(device)
|
| 177 |
+
|
| 178 |
+
for i in reversed(range(len(timestep))):
|
| 179 |
+
t_current =timestep[i]
|
| 180 |
+
t_preq = timestep[i-1] if i > 0 else -1
|
| 181 |
+
t_tensor = torch.full((batch_size,), t_current, dtype=torch.long, device=device)
|
| 182 |
+
|
| 183 |
+
pred_noise = model(img , t_tensor)
|
| 184 |
+
alphas_bar_t = alphas_cumprod_device[t_current]
|
| 185 |
+
alphas_bar_t_preq = alphas_cumprod_device[t_preq] if t_preq >= 0 else torch.tensor(1.0, device=device)
|
| 186 |
+
pred_x0 = (img - torch.sqrt(1 - alphas_bar_t) * pred_noise)/(torch.sqrt(alphas_bar_t))
|
| 187 |
+
pred_x0 = torch.clamp(pred_x0, -1.0, 1.0)
|
| 188 |
+
pred_dir = torch.sqrt(1 - alphas_bar_t_preq) * pred_noise
|
| 189 |
+
img = torch.sqrt(alphas_bar_t_preq) * pred_x0 + pred_dir
|
| 190 |
+
|
| 191 |
+
return img
|