PyTorch中grad()函数工作原理及求导时sum()的作用疑问
PyTorch autograd.grad函数输入参数疑问解答
问题1:为什么需要传入标量的Q.sum()作为grad的第一个参数?Q.sum()的作用是什么?
- PyTorch的
torch.autograd.grad默认设计用于计算标量输出对输入张量的梯度,这是因为深度学习训练场景中几乎所有损失函数都是标量,框架基于这个通用场景做了默认适配。如果第一个参数直接传入非标量的张量,函数无法直接返回符合预期的梯度结果。 - 你的场景中Q是和a形状完全一致的张量,且Q的每个位置元素仅和a对应位置的元素有关。
Q.sum()的作用是将张量Q聚合为单个标量,此时对标量求a的梯度时,sum操作的求导特性刚好可以保留每个位置∂Q_i/∂a_i的结果:求和后的标量对a_i求导时,只有Q中第i个元素的项会产生非零梯度,其余项导数均为0,刚好匹配你需要的逐元素偏导计算需求。
你得到的结果也完全符合数学推导:一阶偏导∂Q/∂a = a²,代入a=[2.,3.]得到[4.,9.];二阶偏导∂²Q/∂a² = 2a,代入得到[4.,6.],和输出完全一致。
问题2:是否可以不使用sum(),直接传入Q作为参数完成导数计算?
可以,此时需要显式传入grad_outputs参数:
当grad的第一个参数为非标量张量时,PyTorch实际计算的是雅可比矩阵和grad_outputs张量的乘积。你只需要传入和Q形状相同的全1张量作为grad_outputs,就能得到和Q.sum()求导完全等价的结果,示例代码如下:
# 无需sum,直接传入Q计算一阶导数 Q_a = torch.autograd.grad(Q, a, grad_outputs=torch.ones_like(Q), create_graph=True)[0] # 二阶导数同理 Q_aa = torch.autograd.grad(Q_a, a, grad_outputs=torch.ones_like(Q_a), create_graph=True)[0]
上述代码运行得到的结果和你原有使用sum的代码结果完全一致。
内容的提问来源于stack exchange,提问作者Prakhar Sharma
相关产品推荐
相关产品推荐

