要完全回答您的问题,需要更长的解释,围绕Backprop 或更根本的chain rule 工作原理的细节展开。
简短的编程答案是Variable 的反向函数计算附加到Variable 的计算图中所有变量的梯度。 (澄清一下:如果你有a = b + c,那么计算图(递归地)首先指向b,然后指向c,然后指向它们的计算方式等)并将这些梯度累积存储(总和)这些变量的.grad 属性。然后,当您调用 opt.step()(即优化器的一个步骤)时,它会将梯度的一部分添加到这些变量的值中。
也就是说,从概念上看,有两个答案:如果您想训练机器学习模型,您通常希望获得关于某个损失函数的梯度。在这种情况下,计算的梯度将使得在应用阶跃函数时整体损失(标量值)将减少。在这种特殊情况下,我们希望将梯度计算为特定值,即单位长度步长(这样学习率将计算出我们想要的梯度分数)。这意味着如果你有一个损失函数,并且你调用loss.backward(),这将与loss.backward(torch.FloatTensor([1.])) 计算相同。
虽然这是 DNN 中反向传播的常见用例,但它只是函数一般微分的一个特例。更一般地,符号微分包(在这种情况下为 autograd,作为 pytorch 的一部分)可用于计算计算图早期部分的梯度,相对于您在任何子图的根处的 any 梯度选择。这是关键字参数gradient 派上用场的时候,因为您可以在那里提供这种“根级”渐变,即使对于非标量函数也是如此!
为了说明,这里有一个小例子:
a = nn.Parameter(torch.FloatTensor([[1, 1], [2, 2]]))
b = nn.Parameter(torch.FloatTensor([[1, 2], [1, 2]]))
c = torch.sum(a - b)
c.backward(None) # could be c.backward(torch.FloatTensor([1.])) for the same result
print(a.grad, b.grad)
打印:
Variable containing:
1 1
1 1
[torch.FloatTensor of size 2x2]
Variable containing:
-1 -1
-1 -1
[torch.FloatTensor of size 2x2]
虽然
a = nn.Parameter(torch.FloatTensor([[1, 1], [2, 2]]))
b = nn.Parameter(torch.FloatTensor([[1, 2], [1, 2]]))
c = torch.sum(a - b)
c.backward(torch.FloatTensor([[1, 2], [3, 4]]))
print(a.grad, b.grad)
打印:
Variable containing:
1 2
3 4
[torch.FloatTensor of size 2x2]
Variable containing:
-1 -2
-3 -4
[torch.FloatTensor of size 2x2]
和
a = nn.Parameter(torch.FloatTensor([[0, 0], [2, 2]]))
b = nn.Parameter(torch.FloatTensor([[1, 2], [1, 2]]))
c = torch.matmul(a, b)
c.backward(torch.FloatTensor([[1, 1], [1, 1]])) # we compute w.r.t. a non-scalar variable, so the gradient supplied cannot be scalar, either!
print(a.grad, b.grad)
打印
Variable containing:
3 3
3 3
[torch.FloatTensor of size 2x2]
Variable containing:
2 2
2 2
[torch.FloatTensor of size 2x2]
和
a = nn.Parameter(torch.FloatTensor([[0, 0], [2, 2]]))
b = nn.Parameter(torch.FloatTensor([[1, 2], [1, 2]]))
c = torch.matmul(a, b)
c.backward(torch.FloatTensor([[1, 2], [3, 4]])) # we compute w.r.t. a non-scalar variable, so the gradient supplied cannot be scalar, either!
print(a.grad, b.grad)
打印:
Variable containing:
5 5
11 11
[torch.FloatTensor of size 2x2]
Variable containing:
6 8
6 8
[torch.FloatTensor of size 2x2]