梯度惩罚反向传播方程推导:如何计算对X'与ε的梯度
判别器输入及混合系数梯度计算方案
前向传播流程
- 步骤1:生成混合批次图像,计算公式为:
X' = ε * real_image + (1 - ε) * generated_image - 步骤2:将
X'传入FFN结构的判别器网络,计算对应输出与损失 - 步骤3:调用深度学习框架内置接口计算判别器网络参数对应的梯度
- 步骤4:计算对输入
X'与混合系数ε的梯度
核心问题
已知FFN结构判别器的参数梯度,如何高效计算对X'与ε的梯度?
Pytorch实现代码
直接使用Pytorch自带的autograd自动求导功能即可实现需求,等效代码如下:
epsilon_shape = [real_data.shape[0]] + [1]*(real_data.dim() - 1) epsilon = torch.rand(epsilon_shape) epsilon = epsilon.to(fake_data.device, fake_data.dtype) real_data = real_data.to(fake_data.dtype) x_hat = epsilon * real_data + (1-epsilon) * fake_data.detach() x_hat.requires_grad = True logits = self.discriminator(x_hat, condition, landmarks) logits = logits.sum() # 计算对x_hat的梯度 grad_x_hat = torch.autograd.grad( outputs=logits, inputs=x_hat, grad_outputs=torch.ones(logits.shape).to(fake_data.dtype).to(fake_data.device), create_graph=True )[0]
ε梯度推导
得到x_hat的梯度后,根据链式法则可直接计算得到ε的梯度:
由x_hat = ε * real_data + (1-ε) * fake_data可推导得导数关系 dx_hat/dε = real_data - fake_data,因此损失对ε的梯度为:grad_epsilon = (grad_x_hat * (real_data - fake_data)).sum(dim=tuple(range(1, real_data.dim())))
维度求和操作是为了匹配ε的广播维度,保证梯度维度和ε一致。
内容的提问来源于stack exchange,提问作者ransomware
相关产品推荐
相关产品推荐

