PyTorch如何记录被覆盖张量以实现自动微分?
PyTorch自动微分如何追踪被覆盖的张量值
PyTorch的自动微分核心依赖计算图,而计算图记录的是张量对象之间的运算依赖关系,和变量名本身没有直接绑定——变量名只是用来指向张量对象的“标签”而已。当你用新张量覆盖变量名时,旧的张量对象并不会消失,它依然存在于已经构建好的计算图中,继续参与后续的微分计算。
我们拆解你给出的代码一步步看:
- 初始化
x(值1.0,需追踪梯度)和w(值2.0,需追踪梯度),执行y = w * x时,PyTorch会在计算图中记录:y是由值为2.0的这个w张量对象和x相乘得到的,同时保存了乘法运算的求导规则。 - 当你执行
w = torch.tensor([4.0], requires_grad=True)时,只是把变量名w的指向改成了一个全新的、值为4.0的张量对象,但之前那个值为2.0的w张量并没有被销毁——它仍然是计算图里y节点的输入之一。 - 后续
z = y ** 2的计算,基于的还是原来的y(关联着旧的w张量),所以z的计算图依然和最初的w=2.0绑定。 - 调用
z.backward()反向传播时,求导路径是:z对y的导数是2y,此时y=2*1=2,所以这部分值为4;y对x的导数是旧w的值2.0;- 最终
x.grad = 4 * 2 = 8,和你得到的结果一致。
简单说:变量名的替换不影响已经建好的计算图,计算图认的是张量对象本身,不是变量名。如果想让新的w参与计算,你需要重新执行y = w * x,让y关联上新的w张量,再构建后续的计算图。
内容的提问来源于stack exchange,提问作者C.M.O.B.
相关产品推荐
相关产品推荐

