:param critic: 判别器模型 :param real: 真实样本 :param fake: 生成的样本 :param device: 设备CUP or GPU :return:
(critic,real,fake,device='cpu')
| 8 | import torch |
| 9 | |
| 10 | def 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 | """ |
| 38 | torch.autograd.grad函数参数如下:https://blog.csdn.net/waitingwinter/article/details/105774720 |