PyTorch LSTMCell调用backward()触发二次反向传播报错排查
报错原因
这个错误的核心是跨训练迭代复用了绑定旧计算图的LSTM隐藏状态,具体触发逻辑:
- 你在训练循环外初始化了
hx、cx,第一次迭代将其传入网络前向传播后,返回的新hx、cx会和第一次迭代的整个计算图建立梯度关联 - PyTorch默认在调用
loss.backward()后,会释放当前反向传播路径上所有保存的中间张量,用来节省显存 - 第二次迭代时,你直接把上一轮返回的、还持有旧计算图引用的
hx、cx传入网络,新生成的前向计算图会连接到已经被释放的旧计算图上。反向传播时需要回溯旧路径上的张量但找不到对应数据,就触发了该报错 - 额外说明:你代码中用
Variable()封装张量的操作是冗余的,PyTorch 0.4版本之后Tensor原生支持自动求导,不需要额外封装,这个操作解决不了当前问题。
修复方案
不要通过添加retain_graph=True的方式修复:这个参数会强制保留所有迭代的计算图,显存占用会随训练步数线性上涨,很快就会出现显存溢出,不符合时序截断反向传播(TBPTT)的训练逻辑。
正确的修复方式是每轮迭代传入隐藏状态前,将其从旧计算图中分离,截断跨迭代的梯度传播,仅保留张量的数值作为下一轮的初始隐藏状态,这也是RNN类模型做截断反向传播的标准写法,既保留了隐藏状态跨步传递的序列记忆能力,又把梯度计算限制在单步迭代内,保证训练效率和显存稳定。具体修改步骤:
- 把训练循环中调用网络的代码行,从原来的
outputs, hx, cx = net(inputs, hx, cx)
修改为
outputs, hx, cx = net(inputs, hx.detach(), cx.detach())
detach()操作会返回一个和原张量数值一致、但和之前所有计算图断开关联的新张量,不会影响隐藏状态的数值传递,同时让每轮的计算图完全独立,反向传播时不会回溯之前迭代的已释放路径。
- (可选优化)删掉所有冗余的
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
相关产品推荐
相关产品推荐

