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这类直接修改原始张量值的操作)修改了计算图中的中间节点,会导致该张量的版本号升高,第二次反向传播时,计算图预期的是修改前的版本号,就会抛出版本不匹配的错误。
解决方案
简化示例修复方案
有两种可选方案:
- 所有batch的loss累积后统一反向传播(最推荐,性能最高,无额外兼容问题)
total_loss = 0 for i in batched_feats: i = i + 5 print(i.shape) total_loss += torch.sum(i) total_loss.backward()
- 多次反向传播时手动控制计算图保留逻辑,最后一次反向传播释放缓存
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)
实际工程代码修复方案
- 优先选择将多批次loss加总后统一反向传播,从根源避免多次反向传播带来的各类问题
- 如果业务逻辑必须多次反向传播,先排查修改代码中所有inplace操作:
- 将
x += a这类自增/自减操作改为x = x + a的非inplace写法 - 将
x[index] = value这类索引赋值操作改为scatter等生成新张量的实现 - 排查自定义算子、第三方依赖调用中是否存在隐式的inplace修改
- 将
内容的提问来源于stack exchange,提问作者TweetysOldFriend
相关产品推荐
相关产品推荐

