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

PyTorch循环中backward()二次反向传播报错及原因咨询

问题:PyTorch循环反向传播二次报错的根因分析

我在使用PyTorch进行模型的循环前向传播与反向传播时,第二次循环触发如下报错:

RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved tensors after calling backward.

排查了网络相关方案,确保单循环内仅执行一次反向传播,且无主动释放计算图的操作,但问题仍未解决。相关代码如下:

criterion = torch.nn.MSELoss().type(data_type) 
optimizer = torch.optim.Adam(net.parameters(), lr=learning_rate) 

image_queue = queue.Queue()
# a separate thread for calculating a value that is used by a main thread
denoise_thread = Thread(target=lambda q, f, p: q.put(f(*p)),  # q -> queue, f -> func, p -> params
                        args=(image_queue, non_local_means,
                              [benchmark_image.clone().squeeze().cpu().detach().numpy(), 3]))
denoise_thread.start()

temp_benchmark = benchmark_image.clone()  #  copy of the benchmark for intermediate calculations
for i in range(num_iter + 1):

    optimizer.zero_grad()

    out = net(net_input)
    temp = benchmark_image - lagrange_multiplier
    loss_net = criterion(out, decrease_image)     # this loss backward normally
    loss_red = criterion(out, temp)               # this is where things go wrong

    total_loss = loss_net + mu * loss_red
    total_loss.backward()                         # FAIL backward!!!

    # updates the benchmark value every iteration a certain number of times
    if i % 30 == 0:
        denoise_thread.join()
        temp_benchmark = image_queue.get()
        temp_benchmark = torch.from_numpy(temp_benchmark)[None, :].cuda()
        temp_benchmark.requires_grad_()
            
        # as before, it is used to update values that will be used by the main thread
        denoise_thread = Thread(target=lambda q, f, p: q.put(f(*p)),
                                args=(image_queue, non_local_means,
                                      [benchmark_image.clone().squeeze().cpu().detach().numpy(), 3]))
        denoise_thread.start()

    benchmark_image = 1 / (beta + mu) * (beta * temp_benchmark + mu * (out + lagrange_multiplier))
    lagrange_multiplier = lagrange_multiplier + out - benchmark_image

    optimizer.step()
       
if denoise_thread.is_alive():
    denoise_thread.join()

经过多次调试,通过修改loss_red = criterion(out, temp.detach_())解决了问题,但仍不清楚报错的根本原因,特此求助。


根因分析

  • 核心问题:计算图的意外复用
    报错本质是第二次反向传播时,尝试复用了第一次循环中已被释放的计算图节点。问题出在benchmark_image和lagrange_multiplier的更新逻辑:

    1. 第一次循环中,temp = benchmark_image - lagrange_multiplier,此时这两个变量还未与模型输出out关联,计算图独立。
    2. 第一次反向传播后,benchmark_image和lagrange_multiplier的更新直接引用了当前循环的out(属于当前计算图的节点):
      benchmark_image = 1 / (beta + mu) * (beta * temp_benchmark + mu * (out + lagrange_multiplier))
      lagrange_multiplier = lagrange_multiplier + out - benchmark_image
      

    这导致第二次循环时,benchmark_image和lagrange_multiplier已包含第一次循环计算图的节点信息。当第二次计算loss_red时,temp依赖的这两个变量关联着已释放的旧计算图,反向传播时回溯该图就触发了报错。

  • 为什么detach()能解决问题?
    temp.detach_()会切断temp与依赖计算图的关联,让loss_red仅基于当前循环的out计算梯度,不再回溯到之前循环中已失效的计算图节点,避免了重复反向传播旧图的问题。

  • 额外优化建议
    代码中temp_benchmark.requires_grad_()是不必要的,因为temp_benchmark是从线程获取的numpy数组转换而来,仅作为基准值参与计算,不需要被反向传播更新,建议移除该设置,避免引入额外的计算图关联。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 20:51:35