为什么PyTorch中backward的RuntimeError出现在循环第二次而非首次?
问题原因解析
这个错误的本质是参数的原地更新和计算图保留逻辑冲突,具体的时序逻辑如下:
第一次循环无报错的原因
- 首次调用
model(x)生成计算图时,使用的是初始版本的所有参数(w1版本号为0),本次生成的loss1和loss2都绑定在同一份计算图上 - 执行
loss1.backward(retain_graph=True)时,仅完成梯度计算、未释放计算图 - 调用
optim.step()原地更新参数,w1版本号变为1,但本次的loss2反向传播时可以直接复用计算图中已存在的中间节点s(已经用初始版w1计算完成),不需要再读取旧版本w1的数值,因此不会触发错误
第二次循环报错的原因
- 再次调用
model(x)生成新计算图时,使用的是版本号为1的w1,本次生成的loss1和loss2都绑定在这份新计算图上,要求依赖的w1版本为1 - 执行
loss1.backward(retain_graph=True)后调用optim.step(),再次原地更新w1,其版本号变为2 - 执行当前迭代的
loss2.backward()时,绑定的计算图要求读取版本号为1的w1,但w1已经被原地修改为版本2,因此触发版本不匹配的RuntimeError
修复方案
方案1:移除不必要的计算图保留
如果没有特殊的跨迭代计算图复用需求,直接删除两个backward调用中的retain_graph=True参数即可,每次反向传播完成后自动释放计算图,不会出现旧计算图引用已更新参数的问题。
方案2:调整参数更新时机
如果有保留单轮迭代内计算图的需求,将两次反向传播放在参数更新之前执行,调整循环内代码顺序如下:
while True: num += 1 print(num) y1_pred, y2_pred = model(x) loss1 = mse(y1_pred, y1) loss2 = mse(y2_pred, y2) optim.zero_grad() loss1.backward(retain_graph=True) loss2.backward() # 第二次反向传播不需要保留计算图 optim.step()
等所有梯度计算完成后再统一更新参数,避免计算图还在使用参数时就被原地修改。
内容的提问来源于stack exchange,提问作者Chongming Gao
相关产品推荐
相关产品推荐

