PyTorch中这段LSTM循环代码是否会破坏计算图?
PyTorch逐步调用LSTM隐状态覆盖问题解答
你担心的变量覆盖导致计算图中断的情况不会发生,也不需要特意把所有中间隐状态存入数组来维持反向传播的依赖关系。
具体原因如下:
- PyTorch的动态计算图是绑定在张量本身的运算历史上的,和变量名的绑定无关。你每次执行
lstm运算得到的新hidden张量,内部已经完整保存了从初始隐状态到当前时间步所有运算的依赖链:上一步的hidden张量作为当前步lstm运算的输入,已经被记录在新hidden的依赖关系里,哪怕你用同一个变量名覆盖了上一步的hidden值,上一步的旧hidden张量因为还被新hidden的依赖链引用,不会被系统回收,反向传播时会自动沿着依赖链回溯所有时间步的运算,不会出现路径中断的问题。 - 简单来说,变量名只是你用来访问张量的临时别名,只要张量本身的运算依赖没有被手动切断(比如调用
.detach()方法),就算别名被覆盖,计算图的依赖链也不会受影响。 - 只有当你的损失计算需要用到每一个时间步的隐状态或者输出时(比如序列标注任务需要拿每一步的输出计算交叉熵),你才需要额外把每一步的
out或者hidden存入列表。如果你的任务只需要用到最后一步的输出或者隐状态,这种逐步覆盖变量的写法是完全合理的。
你提到的代码写法是正确的,符合PyTorch的计算图规则:
for i in inputs: # 逐元素遍历序列,每步执行后hidden存储当前隐状态 out, hidden = lstm(i.view(1, 1, -1), hidden)
内容的提问来源于stack exchange,提问作者Henry Chinner
相关产品推荐
相关产品推荐

