含共享损失项的双PyTorch神经网络反向传播错误解决
解决PyTorch双模型共享损失项的反向传播错误
问题场景
训练两个独立的PyTorch模型M1、M2:
- M1接收输入I1,输出O1;M2接收输入I2,输出O2
- 损失函数定义:
- L1 = some_func(O1, O2) + other_func(O1)
- L2 = some_func(O1, O2) + other_func(O2)
两个损失共享some_func(O1, O2)项
原代码执行时先后触发两个错误:
- 首次报错:
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.
- 设置
retain_graph=True后二次报错:
RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor [64, 1]], which is output 0 of AsStridedBackward0, is at version 2; expected version 1 instead. Hint: the backtrace further above shows the operation that failed to compute its gradient. The variable in question was changed in there or anywhere later. Good luck!
错误原因
- 第一个错误:第一次调用
backward()后,PyTorch默认会释放计算图的中间张量,第二次反向传播时无法复用已释放的计算图,因此需要retain_graph=True保留计算图。 - 第二个错误:原代码先对M2执行
optimizer.step(),直接inplace修改了M2的参数,而后续计算L1的反向传播仍依赖计算图中M2参数的旧版本,导致参数版本不匹配,触发报错。
正确解决方案
核心原则:所有反向传播操作必须在参数更新前完成,避免修改计算图依赖的参数。具体步骤:
- 先计算两个损失函数;
- 清空两个优化器的梯度缓存;
- 依次对两个损失执行反向传播(第一个反向传播需保留计算图);
- 最后分别更新两个模型的参数。
修正后的代码示例
# 1. 先计算两个损失(确保O1、O2的计算图未被破坏) loss_M1 = some_func(O1, O2) + other_func(O1) loss_M2 = some_func(O1, O2) + other_func(O2) # 2. 清空两个优化器的梯度 self.optimizer_M1.zero_grad() self.optimizer_M2.zero_grad() # 3. 反向传播:先处理M1的损失,保留计算图供M2使用 loss_M1.backward(retain_graph=True) # 处理M2的损失,此时计算图可释放(默认retain_graph=False) loss_M2.backward() # 4. 分别更新两个模型的参数 self.optimizer_M1.step() self.optimizer_M2.step()
补充说明
如果你的计算图非常庞大,retain_graph=True会占用更多显存,此时可以改用torch.autograd.grad手动计算并累加梯度,避免保留整个计算图:
# 计算M1的梯度 grads_M1 = torch.autograd.grad(loss_M1, self.M1.parameters(), retain_graph=True) # 计算M2的梯度 grads_M2 = torch.autograd.grad(loss_M2, self.M2.parameters()) # 手动为参数赋值梯度 for param, grad in zip(self.M1.parameters(), grads_M1): param.grad = grad for param, grad in zip(self.M2.parameters(), grads_M2): param.grad = grad # 更新参数 self.optimizer_M1.step() self.optimizer_M2.step()
内容的提问来源于stack exchange,提问作者Manav Mishra
相关产品推荐
相关产品推荐

