adityachaubey commited on
Commit
9122c78
·
1 Parent(s): 9596f34

Add DDIM sampling function for image generation

Browse files
Files changed (1) hide show
  1. model.py +24 -1
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