MCPcopy Create free account
hub / github.com/KeepTryingTo/Pytorch-GAN / gradient_penality

Function gradient_penality

WGANGPCode/utils.py:10–35  ·  view source on GitHub ↗

:param critic: 判别器模型 :param real: 真实样本 :param fake: 生成的样本 :param device: 设备CUP or GPU :return:

(critic,real,fake,device='cpu')

Source from the content-addressed store, hash-verified

8import torch
9
10def gradient_penality(critic,real,fake,device='cpu'):
11 """
12 :param critic: 判别器模型
13 :param real: 真实样本
14 :param fake: 生成的样本
15 :param device: 设备CUP or GPU
16 :return:
17 """
18 BATCH_SIZE,C,H,W = real.shape
19 alpha = torch.randn(size=(BATCH_SIZE,1,1,1)).repeat(1,C,H,W).to(device)
20 interpolated_images=real*alpha + fake*(1-alpha)
21
22 #计算判别器输出
23 mixed_scores = critic(interpolated_images)
24 #求导
25 gradient = torch.autograd.grad(
26 inputs=interpolated_images,
27 outputs=mixed_scores,
28 grad_outputs=torch.ones_like(mixed_scores),
29 create_graph=True,
30 retain_graph=True
31 )[0]
32 gradient = gradient.view(gradient.shape[0],-1)
33 gradient_norm = gradient.norm(2,dim = 1)
34 gradient_penality = torch.mean((gradient_norm - 1)**2)
35 return gradient_penality
36
37"""
38torch.autograd.grad函数参数如下:https://blog.csdn.net/waitingwinter/article/details/105774720

Callers 1

train.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected