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

PyTorch LSTMCell调用backward()触发二次反向传播报错排查

报错原因

这个错误的核心是跨训练迭代复用了绑定旧计算图的LSTM隐藏状态,具体触发逻辑:

  • 你在训练循环外初始化了hx、cx,第一次迭代将其传入网络前向传播后,返回的新hx、cx会和第一次迭代的整个计算图建立梯度关联
  • PyTorch默认在调用loss.backward()后,会释放当前反向传播路径上所有保存的中间张量,用来节省显存
  • 第二次迭代时,你直接把上一轮返回的、还持有旧计算图引用的hx、cx传入网络,新生成的前向计算图会连接到已经被释放的旧计算图上。反向传播时需要回溯旧路径上的张量但找不到对应数据,就触发了该报错
  • 额外说明:你代码中用Variable()封装张量的操作是冗余的,PyTorch 0.4版本之后Tensor原生支持自动求导,不需要额外封装,这个操作解决不了当前问题。
修复方案

不要通过添加retain_graph=True的方式修复:这个参数会强制保留所有迭代的计算图,显存占用会随训练步数线性上涨,很快就会出现显存溢出,不符合时序截断反向传播(TBPTT)的训练逻辑。

正确的修复方式是每轮迭代传入隐藏状态前,将其从旧计算图中分离,截断跨迭代的梯度传播,仅保留张量的数值作为下一轮的初始隐藏状态,这也是RNN类模型做截断反向传播的标准写法,既保留了隐藏状态跨步传递的序列记忆能力,又把梯度计算限制在单步迭代内,保证训练效率和显存稳定。具体修改步骤:

  1. 把训练循环中调用网络的代码行,从原来的
outputs, hx, cx  = net(inputs, hx, cx)

修改为

outputs, hx, cx  = net(inputs, hx.detach(), cx.detach())

detach()操作会返回一个和原张量数值一致、但和之前所有计算图断开关联的新张量,不会影响隐藏状态的数值传递,同时让每轮的计算图完全独立,反向传播时不会回溯之前迭代的已释放路径。

  1. (可选优化)删掉所有冗余的Variable()封装,直接创建Tensor即可,对应代码修改示例:
# 初始化隐藏状态,移除Variable封装
hx = torch.randn(3, 10)
cx = torch.randn(3, 10)

# 训练循环内创建输入、标签,移除Variable封装
inputs = torch.randn(10, 3, 10) # time step, batch,hidden size 
labels = torch.randn(10, 3, 10)

修改完成后代码可以正常跑通训练,不会再触发该反向传播错误,同时显存占用保持稳定。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 21:57:26