【问题标题】:Explanation of the tuple grad_fn.next_functions in PyTorchPyTorch中元组grad_fn.next_functions的解释
【发布时间】:2020-12-30 11:24:43
【问题描述】:
火炬计算图由grad_fn组成。对于终止节点,grad_fn 对象有一个名为next_functions 的属性,它是一个元组。我知道使用元组的第一个元素(第 0 个索引),我可以重建梯度的计算图。但我想知道元组的第二个元素(第一个索引)是什么意思?
在 PyTorch 论坛的one of the answers 中,据说:
数字是下一个后向函数的输入数字,因此只能在函数具有多个可微输出时为非零(没有那么多,但例如 RNN 函数通常这样做)。
但我不明白这个说法。有人可以举个例子来解释一下吗?
【问题讨论】:
标签:
python
pytorch
autograd
【解决方案1】:
我不是 Pytorch 方面的专家,但试图从示例中回答您的问题:
a, b = torch.randn(2, requires_grad=True).unbind()
c = a+b
print(c.grad_fn.next_functions)
>>> ((<UnbindBackward object at 0x7f0ea438de80>, 0), (<UnbindBackward object at 0x7f0ea438de80>, 1))
现在,将任何 Pytorch 函数视为生成“输出列表”,而不是“输出”。所以,如果一个函数只产生一个输出(典型情况);它生成一个等于[output] 的列表。但是,如果函数产生多个输出,则它有一个 len > 1 个输出的列表,例如:[output0, output1]。
由此,我理解元组成分如下:
(grad_fn: the function object that resulted in this tensor, i: index of the tensor in the function outputs list.. which is typically zero since functions typically have one output)
将这种理解应用于代码,unbind 函数有两个输出:a 位于输出“列表”的索引 0,b 位于输出“列表”的索引 1。通图有以下几点推理:
-
c.grad_fn 是一个 AddBackward 对象,因为 c 是加法运算的结果。该加法运算在计算图中有两个分支(因为它添加了两个操作数,a 和 b)。第一个分支用于a,第二个分支用于b(一个分支用于操作数,按顺序排列)。
- 可以说
output_tuple = c.grad_fun.next_functions,output_tuple 有 2 个元素,output_tuple[0] 是 a 的(grad_fn,a 的索引在取消绑定输出中),output_tuple[1] 是 b' s(grad_fn,解除绑定输出中b的索引)。
-
a 的 grad_fn 和 b 的 grad_fn 相同(即完全相同的对象,a.grad_fn == b.grad_fn 返回 True)。但是,a 的元组有第二个元素 = 0,因为它是 unbind 函数的 第一个 输出,b 的元组有第二个元素 = 1,因为它是 unbind 函数的第二个输出。
即grad_fn 元组中的第二个条目是输出的生产函数列表中张量的索引。