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 = 2d(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
相关产品推荐
相关产品推荐

