PyTorch中retain_graph作用及两次backward无报错原因问询
PyTorch两次反向传播报错差异的原因解析
运行第一段代码时,若不在y1.backward()中传入retain_graph=True,执行y2.backward()会触发RuntimeError;但运行第二段代码时,y1和y2共享节点z,执行y1.backward()后再运行y2.backward()却无报错,这两者的差异原因是什么?
第一段代码
import torch x = torch.tensor([2.0], requires_grad=True) y = torch.tensor([3.0], requires_grad=True) f = x+y z = 2*f y1 = z**2 y2 = z**3 y1.backward() y2.backward()
报错信息
Traceback (most recent call last): File "/Users/a0m08er/pytorch/pytorch_tutorial/tensor.py", line 58, in <module> y2.backward() File "/Users/a0m08er/pytorch/lib/python3.11/site-packages/torch/_tensor.py", line 521, in backward torch.autograd.backward( File "/Users/a0m08er/pytorch/lib/python3.11/site-packages/torch/autograd/__init__.py", line 289, in backward _engine_run_backward( File "/Users/a0m08er/pytorch/lib/python3.11/site-packages/torch/autograd/graph.py", line 769, in _engine_run_backward return Variable._execution_engine.run_backward( # Calls into the C++ engine to run the backward pass ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved tensors after calling backward.
第二段代码
import torch x = torch.tensor([2.0], requires_grad=True) y = torch.tensor([3.0], requires_grad=True) z = x+y y1 = z**2 y2 = z**3 y1.backward() y2.backward()
差异原因
核心在于PyTorch反向传播时计算图的资源释放逻辑——反向传播默认会释放不再被依赖的节点的计算图资源,具体差异如下:
第一段代码的计算图路径:
x/y → f → z → y1和x/y → f → z → y2
执行y1.backward()时,反向传播从y1回溯到z,再到f,最后到x和y。此时f节点的资源会被释放,因为它只被z依赖,而z的反向传播完成后,f没有其他未完成反向传播的节点依赖它。当执行y2.backward()时,需要再次回溯到z→f→x/y,但f的资源已经被释放,因此触发报错。第二段代码的计算图路径:
x/y → z → y1和x/y → z → y2
执行y1.backward()时,反向传播从y1回溯到z,再到x和y。此时z节点不会被释放,因为它还被y2依赖(PyTorch会追踪节点的依赖计数)。所以执行y2.backward()时,依然可以正常回溯到z→x/y,不会触发报错。
简单总结:第一段代码中f节点在第一次反向传播后无后续依赖,被释放;第二段代码中z节点仍被y2依赖,资源保留,第二次反向传播可正常进行。
内容的提问来源于stack exchange,提问作者pasternak
相关产品推荐
相关产品推荐

