PyTorch循环中backward()二次反向传播报错及原因咨询
我在使用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的更新逻辑:- 第一次循环中,
temp = benchmark_image - lagrange_multiplier,此时这两个变量还未与模型输出out关联,计算图独立。 - 第一次反向传播后,
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

