PyTorch:修复使用中间变量梯度参与后续计算的报错问题
问题根源
你遇到的错误本质是:
.grad存储的是梯度的数值结果,不属于计算图的一部分,用它构建的新张量f没法回溯到原始输入(x_verts、y_verts等)的计算路径,自然无法完成二次反向传播。- 第一次调用
z.backward()后,PyTorch默认会释放计算图,后续再基于这个图的衍生操作会报错。
修复方案
不要直接提取.grad,而是用torch.autograd.grad函数直接计算梯度(返回的是带计算图的张量),保留梯度与原始输入的关联,这样就能完成二次反向传播。
修改后的代码如下:
import torch device = torch.device("cuda" if torch.cuda.is_available() else "cpu") point_2 = torch.tensor([0.2, 0.8], device=device, requires_grad=True) p = torch.cat((point_2, torch.tensor([0], device=device)), 0) x_verts = torch.tensor([0.0, 1.0, 0.0], device=device, requires_grad=True) y_verts = torch.tensor([0.0, 0.0, 1.0], device=device, requires_grad=True) z_verts = torch.tensor([0.1, -0.1, 0.2], device=device, requires_grad=True) v1_2d = torch.cat((torch.index_select(x_verts, 0, torch.tensor([0], device=device)), torch.index_select(y_verts, 0, torch.tensor([0], device=device)), torch.tensor([0], device=device))) v2_2d = torch.cat((torch.index_select(x_verts, 0, torch.tensor([1], device=device)), torch.index_select(y_verts, 0, torch.tensor([1], device=device)), torch.tensor([0], device=device))) v3_2d = torch.cat((torch.index_select(x_verts, 0, torch.tensor([2], device=device)), torch.index_select(y_verts, 0, torch.tensor([2], device=device)), torch.tensor([0], device=device))) area_3 = torch.cross(v2_2d - v1_2d, v3_2d - v1_2d) area = torch.index_select(area_3, 0, torch.tensor([2], device=device)) alpha_3 = 0.5 * torch.cross(v2_2d - p, v3_2d - p) / area beta_3 = 0.5 * torch.cross(v3_2d - p, v1_2d - p) / area gamma_3 = 0.5 * torch.cross(v1_2d - p, v2_2d - p) / area alpha = torch.index_select(alpha_3, 0, torch.tensor([2], device=device)) beta = torch.index_select(beta_3, 0, torch.tensor([2], device=device)) gamma = torch.index_select(gamma_3, 0, torch.tensor([2], device=device)) z = alpha * torch.index_select(z_verts, 0, torch.tensor([0], device=device)) + \ beta * torch.index_select(z_verts, 0, torch.tensor([1], device=device)) + \ gamma * torch.index_select(z_verts, 0, torch.tensor([2], device=device)) # 关键修改:用autograd.grad计算z对point_2的梯度,返回的是带计算图的张量 grad_point2, = torch.autograd.grad(z, point_2, retain_graph=True) grad_norm = torch.norm(grad_point2) f = torch.tanh(10.0 * (grad_norm - 2.0)) f.backward() # 现在可以正常反向传播到原始参数 print(x_verts.grad) print(y_verts.grad) print(z_verts.grad)
关键修改点说明
- 用
torch.autograd.grad(z, point_2, retain_graph=True)替代直接取point_2.grad:- 这个函数返回的是
z对point_2的梯度张量,保留了与原始计算图的关联,后续用它计算f时,梯度可以正常回溯到x_verts、y_verts等参数。 retain_graph=True确保第一次计算梯度后不释放计算图,因为后续还要对f进行反向传播。
- 这个函数返回的是
- 移除了原来的
z.backward(),因为torch.autograd.grad已经完成了一次梯度计算,同时保留了计算图供后续使用。 - 补充了
device的定义(原代码里未显式声明,避免运行报错),并且所有张量统一指定了device,防止CPU/GPU张量混合的问题。
内容的提问来源于stack exchange,提问作者Cedric Martens
相关产品推荐
相关产品推荐

