为何出现RuntimeError: Trying to backward through the graph a second time错误?
PyTorch RuntimeError 问题排查与解决
问题描述
运行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.
核心疑问:
- 程序为何尝试对同一计算图执行第二次反向传播?
- 是否存在直接访问已释放保存张量的情况?
代码如下:
import torch import random image_width, image_height = 128, 128 def apply_ellipse_mask(img, pos, axes): r = torch.arange(image_height)[:, None] c = torch.arange(image_width)[None, :] val_array = ((c - pos[0]) ** 2) / axes[0] ** 2 + ((r - pos[1]) ** 2) / axes[1] ** 2 mask = torch.where((0.9 < val_array) & (val_array < 1), torch.tensor(1.0), torch.tensor(0.0)) return img * (1.0 - mask) + mask random.seed(0xced) sphere_radius = image_height / 3 sphere_position = torch.tensor([image_width / 2, image_height / 2 ,0], requires_grad=True) ref_image = apply_ellipse_mask(torch.zeros(image_width, image_height, requires_grad=True), sphere_position, [sphere_radius, sphere_radius, sphere_radius]) ellipsoid_pos = torch.tensor([sphere_position[0], sphere_position[1], 0], requires_grad=True) ellipsoid_axes = torch.tensor([image_width / 3 + (random.random() - 0.5) * image_width / 5, image_height / 3 + (random.random() - 0.5) * image_height / 5, image_height / 2], requires_grad=True) optimizer = torch.optim.Adam([ellipsoid_axes], lr=0.1) criterion = torch.nn.MSELoss() for _ in range(100): optimizer.zero_grad() current_image = torch.zeros(image_width, image_height, requires_grad=True) current_image = apply_ellipse_mask(current_image, ellipsoid_pos, ellipsoid_axes) loss = criterion(current_image, ref_image) loss.backward() print(_, loss) optimizer.step()
问题原因分析
ref_image计算图未断开ref_image是基于带梯度追踪的张量(sphere_position、torch.zeros(..., requires_grad=True))生成的,其计算图一直被保留。每次迭代计算损失时,ref_image的计算图会被纳入当前反向传播链路。第一次反向传播后,ref_image相关的中间张量已被释放,第二次迭代再次触发反向传播时,就会因访问已释放的张量而报错。不必要的梯度追踪设置
ref_image作为固定参考图,不需要参与梯度更新,却被开启了梯度追踪。- 迭代中创建的
current_image无需单独设置requires_grad=True,因为待优化参数ellipsoid_axes已开启梯度追踪,初始零张量的梯度追踪是多余的。
解决方法
方案一:断开ref_image的计算图
通过两种方式让ref_image脱离梯度计算链路:
- 使用
detach():ref_image = apply_ellipse_mask(...).detach() - 使用
torch.no_grad()上下文管理器包裹计算过程:
with torch.no_grad(): ref_image = apply_ellipse_mask(torch.zeros(image_width, image_height), sphere_position, [sphere_radius, sphere_radius, sphere_radius])
方案二:移除冗余的requires_grad=True
sphere_position是固定参考位置,创建时设requires_grad=False;- 迭代中创建
current_image时,直接用torch.zeros(image_width, image_height),无需开启梯度追踪。
修改后的完整代码
import torch import random image_width, image_height = 128, 128 def apply_ellipse_mask(img, pos, axes): r = torch.arange(image_height)[:, None] c = torch.arange(image_width)[None, :] val_array = ((c - pos[0]) ** 2) / axes[0] ** 2 + ((r - pos[1]) ** 2) / axes[1] ** 2 mask = torch.where((0.9 < val_array) & (val_array < 1), torch.tensor(1.0), torch.tensor(0.0)) return img * (1.0 - mask) + mask random.seed(0xced) sphere_radius = image_height / 3 # 参考位置无需梯度追踪 sphere_position = torch.tensor([image_width / 2, image_height / 2 ,0], requires_grad=False) # 上下文管理器断开参考图的计算图 with torch.no_grad(): ref_image = apply_ellipse_mask(torch.zeros(image_width, image_height), sphere_position, [sphere_radius, sphere_radius, sphere_radius]) ellipsoid_pos = torch.tensor([sphere_position[0], sphere_position[1], 0], requires_grad=True) ellipsoid_axes = torch.tensor([image_width / 3 + (random.random() - 0.5) * image_width / 5, image_height / 3 + (random.random() - 0.5) * image_height / 5, image_height / 2], requires_grad=True) optimizer = torch.optim.Adam([ellipsoid_axes], lr=0.1) criterion = torch.nn.MSELoss() for _ in range(100): optimizer.zero_grad() # 初始零张量无需梯度追踪 current_image = torch.zeros(image_width, image_height) current_image = apply_ellipse_mask(current_image, ellipsoid_pos, ellipsoid_axes) loss = criterion(current_image, ref_image) loss.backward() print(_, loss.item()) optimizer.step()
内容的提问来源于stack exchange,提问作者Cedric Martens
相关产品推荐
相关产品推荐

