x_start = x_0
(self, model: nn.Module, x_start: Tensor, x_noise: Tensor = None, model_kwargs=None)
| 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) |