为何torch.autograd.Function的backward方法需传入grad_output?
为什么torch.autograd.Function的backward方法需要grad_output参数?
这个问题的核心是链式法则——PyTorch的自动微分系统就是靠链式法则把损失的梯度一步步传递回输入张量的。
简单来说:
grad_output是上游传来的梯度,也就是最终损失值L对当前自定义函数输出y的梯度,即dL/dy。- 你在backward里需要计算的是损失对当前函数输入
x的梯度dL/dx。根据链式法则,dL/dx = dL/dy * dy/dx,其中dy/dx是你自定义函数的局部梯度(输出对输入的导数)。
拿你提供的LegendrePolynomial3例子来说:
- forward方法实现的是
y = 0.5*(5x³ - 3x),它的局部导数dy/dx = 1.5*(5x² - 1)。 - backward方法接收的
grad_output就是dL/dy,把它和局部导数相乘,就得到了dL/dx,这正是我们需要返回给上游的梯度。
为什么必须传这个参数?因为你的自定义函数几乎不会是计算图的最后一步——它的输出会被后续的层或操作继续处理,最终才得到损失。只有拿到上游传来的dL/dy,才能把梯度正确传递回输入x。如果没有这个参数,你只能算出局部的dy/dx,但无法和损失挂钩,也就完成不了完整的反向传播。
你已进行的尝试:
- 阅读PyTorch官方教程《Learning PyTorch with Examples》
该教程讲解了如何定义自定义autograd函数,你理解LegendrePolynomial3类的forward方法实现为½ * (5x³ - 3x),但仍不清楚backward方法为何需要grad_output参数。
相关代码示例:
class LegendrePolynomial3(torch.autograd.Function): """ 我们可以通过继承torch.autograd.Function并实现前向和反向传播(基于张量操作), 来定义自己的自定义autograd函数。 """ @staticmethod def forward(ctx, input): """ 在前向传播中,我们接收一个包含输入的张量并返回一个包含输出的张量。 ctx是上下文对象,可用于存储反向计算所需的信息。你可以使用ctx.save_for_backward方法 缓存任意对象供反向传播使用。 """ ctx.save_for_backward(input) # ½ * (5x³ - 3x) return 0.5 * (5 * input ** 3 - 3 * input) @staticmethod def backward(ctx, grad_output): """ 在反向传播中,我们接收一个包含损失对输出的梯度的张量, 需要计算损失对输入的梯度。 """ input, = ctx.saved_tensors # d/dx ½ * (5x³ - 3x) # d/dx (½ * 5x³) - (½ * 3x) # (3 * ½ * 5x²) - (1 * ½ * 3) # 1.5 * (5x² - 1) return grad_output * 1.5 * (5 * input ** 2 - 1)
内容的提问来源于stack exchange,提问作者Jason Rich Darmawan
相关产品推荐
相关产品推荐

