| 287 | return self._loss_d_real + self._loss_d_cond + self._loss_d_fake, fake_imgs_masked |
| 288 | |
| 289 | def _gradinet_penalty_D(self, fake_imgs_masked): |
| 290 | # interpolate sample |
| 291 | alpha = torch.rand(self._B, 1, 1, 1).cuda().expand_as(self._real_img) |
| 292 | interpolated = Variable(alpha * self._real_img.data + (1 - alpha) * fake_imgs_masked.data, requires_grad=True) |
| 293 | interpolated_prob, _ = self._D(interpolated) |
| 294 | |
| 295 | # compute gradients |
| 296 | grad = torch.autograd.grad(outputs=interpolated_prob, |
| 297 | inputs=interpolated, |
| 298 | grad_outputs=torch.ones(interpolated_prob.size()).cuda(), |
| 299 | retain_graph=True, |
| 300 | create_graph=True, |
| 301 | only_inputs=True)[0] |
| 302 | |
| 303 | # penalize gradients |
| 304 | grad = grad.view(grad.size(0), -1) |
| 305 | grad_l2norm = torch.sqrt(torch.sum(grad ** 2, dim=1)) |
| 306 | self._loss_d_gp = torch.mean((grad_l2norm - 1) ** 2) * self._opt.lambda_D_gp |
| 307 | |
| 308 | return self._loss_d_gp |
| 309 | |
| 310 | def _compute_loss_D(self, estim, is_real): |
| 311 | return -torch.mean(estim) if is_real else torch.mean(estim) |