MCPcopy Create free account
hub / github.com/albertpumarola/GANimation / _gradinet_penalty_D

Method _gradinet_penalty_D

models/ganimation.py:289–308  ·  view source on GitHub ↗
(self, fake_imgs_masked)

Source from the content-addressed store, hash-verified

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)

Callers 1

optimize_parametersMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected