为何同一PyTorch模型下两个反向传播案例一个报错一个正常?
PyTorch反向传播计算图疑问解析
前置条件
import torch.nn as nn class Net(nn.Module): def __init__(self, input, H, output): super(Net, self).__init__() self.linear1 = nn.Linear(input, H) self.linear2 = nn.Linear(H, output) def forward(self, x): x = torch.sigmoid(self.linear1(x)) x = self.linear2(x) return x
案例1:无报错
net = Net(2, 3, 2) input = torch.randn(2) loss1 = net(input).sum() loss2 = net(input).sum()/2 loss1.backward() loss2.backward()
案例2:触发错误
报错信息:RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed)
torch.autograd.set_detect_anomaly(True) net = Net(2, 3, 2) input = torch.randn(2) loss1 = net(input).sum() loss2 = loss1/2 loss1.backward() loss2.backward()
问题
我已查阅过类似问题的解答,了解到报错原因是第一次反向传播后计算图被丢弃。但按照这个逻辑,案例1中两个loss共享以net为根节点的子计算图,第一次反向传播后计算图应该被丢弃,第二次反向传播也应报错,但实际并未报错,请问这是为什么?
原因解析
核心区别在于两次反向传播对应的是完全独立的计算图,而非共享同一张图:
- 案例1中,
loss1和loss2分别来自两次独立的net(input)调用。每次调用net(input)都会从头构建全新的计算图,包含从input到线性层、sigmoid再到输出的完整链路。因此loss1.backward()只会释放它自己对应的那张计算图,loss2对应的是另一张完全独立的图,反向传播时自然不会有问题。 - 案例2中,
loss2是基于loss1计算得到的,两者共享同一张从input到loss1的计算图。第一次loss1.backward()执行后,PyTorch默认会销毁这张图的中间缓存(节省内存),当再对loss2执行反向传播时,需要访问已经被释放的图结构,就会触发报错。
简单来说,案例1是两张独立的图各自完成一次反向传播,案例2是同一张图尝试两次反向传播,这就是两者结果不同的根本原因。
内容的提问来源于stack exchange,提问作者zhixin
相关产品推荐
相关产品推荐

