You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中retain_grad()位置影响梯度结果的原因探究

PyTorch中retain_grad()调用位置导致梯度差异的原因

这两段代码的核心差异在于**retain_grad()的调用时机对计算图追踪逻辑的影响**,以及PyTorch对in-place操作的处理规则:

第一种代码逻辑(正常梯度输出)

import torch
a = torch.tensor([1.], requires_grad=True)
y = torch.zeros((10))
gt = torch.zeros((10))

y[0] = a
y[1] = y[0] * 2
y.retain_grad()

loss = torch.sum((y-gt) ** 2)
loss.backward()
print(y.grad)
  • 执行流程:先完成对y的所有元素赋值操作,y[0]=a和y[1]=y[0]*2都会被PyTorch追踪并构建完整计算图——y1依赖y0,y0依赖a。
  • 调用y.retain_grad()后,y被标记为需要保留梯度的非叶子节点。反向传播时,PyTorch会按照计算图逐一计算每个元素的梯度:
    • d(loss)/dy0 = 2*(y0 - gt) = 2*1 = 2
    • d(loss)/dy1 = 2*(y1 - gt) = 2*2 = 4
    • 其余元素与a无依赖关系,梯度为0,最终输出[2.,4.,0.,...]。

第二种代码逻辑(梯度异常)

import torch
a = torch.tensor([1.], requires_grad=True)
y = torch.zeros((10))
gt = torch.zeros((10))

y[0] = a
y.retain_grad()
y[1] = y[0] * 2

loss = torch.sum((y-gt) ** 2)
loss.backward()
print(y.grad)
  • 执行流程:先调用y.retain_grad(),此时y被标记为需要保留梯度的节点,PyTorch会停止追踪后续对y的in-place修改操作。
  • y[1] = y[0]*2虽然将y1的数值设为2a,但这个操作没有被记录到计算图中,y1与a不存在依赖关系。
  • 反向传播时,只有y0与a存在直接依赖,因此y0的梯度等于a的梯度(d(loss)/da = 10*1 =10),y1及其他元素因无计算图依赖,梯度为0,最终输出[10.,0.,0.,...]。

关键结论

  • retain_grad()会改变PyTorch对目标张量的追踪规则:调用后,对该张量的in-place修改不会被纳入计算图。
  • 若要保留完整的计算图依赖,需在所有张量操作完成后再调用retain_grad()。

内容的提问来源于stack exchange,提问作者Decarbonized formaldehyde

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 19:35:20