MCPcopy Create free account
hub / github.com/CompVis/diff2flow / training_losses

Method training_losses

diff2flow/ddpm.py:218–262  ·  view source on GitHub ↗

x_start = x_0

(self, model: nn.Module, x_start: Tensor, x_noise: Tensor = None, model_kwargs=None)

Source from the content-addressed store, hash-verified

216 )
217
218 def training_losses(self, model: nn.Module, x_start: Tensor, x_noise: Tensor = None, model_kwargs=None):
219 """ x_start = x_0 """
220 model_kwargs = model_kwargs or {}
221 bs, dev = x_start.shape[0], x_start.device
222
223 # uniformly sample t
224 t = torch.randint(0, self.num_timesteps, (bs,), device=dev).long()
225
226 # sample xt ~ q(xt | x0)
227 noise = torch.randn_like(x_start) if x_noise is None else x_noise
228 x_t = self.q_sample(x_start=x_start, t=t, noise=noise)
229
230 model_out = model(x_t, t, **model_kwargs)
231
232 if self.parameterization == "eps":
233 target = noise
234 elif self.parameterization == "x0":
235 target = x_start
236 elif self.parameterization == "v":
237 target = self.get_v(x_start, noise, t)
238 else:
239 raise NotImplementedError(f"Parameterization {self.parameterization} not yet supported")
240
241 # compute loss
242 if self.loss_type == "l1":
243 loss = (target - model_out).abs()
244 elif self.loss_type == "l2":
245 loss = torch.nn.functional.mse_loss(target, model_out)
246 else:
247 raise NotImplementedError(f"Loss type {self.loss_type} not yet supported")
248 loss = loss.mean(dim=[*range(1, loss.ndim)]) # (bs,)
249
250 # losses
251 loss_dict = {}
252
253 loss_dict['loss_simple'] = loss.mean()
254 loss_simple = loss.mean() * self.l_simple_weight
255
256 loss_vlb = (self.lvlb_weights[t] * loss).mean()
257 loss_dict['loss_vlb'] = loss_vlb
258
259 loss = loss_simple + self.original_elbo_weight * loss_vlb
260 loss_dict['loss'] = loss
261
262 return loss, loss_dict
263
264 def forward(self, *args, **kwargs):
265 return self.training_losses(*args, **kwargs)

Callers 1

forwardMethod · 0.95

Calls 2

q_sampleMethod · 0.95
get_vMethod · 0.95

Tested by

no test coverage detected