You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

梯度惩罚反向传播方程推导:如何计算对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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.02 08:45:04