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

PyTorch张量切片时backward触发二次反向传播RuntimeError问题求解

简化示例报错的根本原因

先确认参数计算结果:输入的vectors长度为10,factor = np.ceil(10/6) = 2,因此batch_matrix会返回2*2=4个从pmatrix切分得到的子张量。所有子张量共享从初始vectors生成pmatrix的完整计算图,包括concat_matrix函数中生成的project_x、project_y等中间节点。

PyTorch默认执行backward()反向传播后,会立即释放计算图所有中间节点的缓存以节省内存。第一次反向传播完成后,上游的中间缓存就已经被清空,后续对其他子张量调用反向传播时,找不到需要的中间计算结果就会抛出错误。

前两次调用不报错、第三次才报错属于PyTorch内存回收机制的偶然表现,是因为前两次反向传播依赖的部分中间结果刚好还没被完全回收,不代表代码逻辑合规。

加retain_graph=True后报inplace错误的原因

retain_graph=True会强制保留所有计算图中间缓存,支持多次反向传播。但如果你的代码中存在inplace操作(比如x += 1、x[mask] = 0这类直接修改原始张量值的操作)修改了计算图中的中间节点,会导致该张量的版本号升高,第二次反向传播时,计算图预期的是修改前的版本号,就会抛出版本不匹配的错误。

解决方案

简化示例修复方案

有两种可选方案:

  1. 所有batch的loss累积后统一反向传播(最推荐,性能最高,无额外兼容问题)
total_loss = 0
for i in batched_feats:
    i = i + 5
    print(i.shape)
    total_loss += torch.sum(i)
total_loss.backward()
  1. 多次反向传播时手动控制计算图保留逻辑,最后一次反向传播释放缓存
for idx, i in enumerate(batched_feats):
    i = i + 5
    print(i.shape)
    summed = torch.sum(i)
    # 最后一次反向传播不需要保留计算图
    summed.backward(retain_graph = idx != len(batched_feats)-1)

实际工程代码修复方案

  1. 优先选择将多批次loss加总后统一反向传播,从根源避免多次反向传播带来的各类问题
  2. 如果业务逻辑必须多次反向传播,先排查修改代码中所有inplace操作:
    • 将x += a这类自增/自减操作改为x = x + a的非inplace写法
    • 将x[index] = value这类索引赋值操作改为scatter等生成新张量的实现
    • 排查自定义算子、第三方依赖调用中是否存在隐式的inplace修改

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 03:45:03