自定义torch.autograd.Function前反向函数异常,求代码问题排查
问题
我想用PyTorch优化一个包含torch.autograd无法追踪操作的目标函数,按照官方指引把操作封装在torch.autograd.Function的自定义forward()方法里,也手动实现了已知梯度的backward()方法,但代码无法正常工作。用torchviz绘制计算图后,发现grad_x和任何节点都没有连接。
自定义Function代码:
class Projector(torch.autograd.Function): # non_torch_var are constant values needed by the operation @staticmethod def forward(ctx, vertices, non_torch_var1, non_torch_var2, non_torch_var3): ctx.save_for_backward(vertices) vertices2=vertices.detach().cpu().numpy() ctx.non_torch_var1 = non_torch_var1 ctx.non_torch_var2 = non_torch_var2 ctx.non_torch_var3 = non_torch_var3 out = project_mesh(vertices2, non_torch_var1, non_torch_var2, non_torch_var3) out = torch.tensor(out, requires_grad=True) return out @staticmethod def backward(ctx, grad_out): vertices = ctx.saved_tensors[0] vertices2 = vertices.detach().cpu().numpy() non_torch_var1 = ctx.non_torch_var1 non_torch_var2 = ctx.non_torch_var2 non_torch_var3 = ctx.non_torch_var3 grad_vertices = grad_project_mesh(vertices2, non_torch_var1, non_torch_var2, non_torch_var3) grad_vertices = torch.tensor(grad_vertices, requires_grad=True) return grad_vertices, None, None, None
绘制计算图的代码:
import torchviz out = Projector.apply(*input) grad_x, = torch.autograd.grad(out.sum(), vertices, create_graph=True) torchviz.make_dot((grad_x, vertices, out), params={"grad_x": grad_x, "vertices": vertices, "out": out}).render("attached", format="png")
问题分析与修复方案
你的代码存在三个关键错误,直接导致计算图关联断裂、梯度计算异常:
错误1:手动设置输出tensor的requires_grad=True
自定义Function的输出与输入的梯度关联由PyTorch autograd框架自动管理,不需要手动给返回的tensor设置requires_grad=True,这会破坏框架对计算图的追踪逻辑,导致输出节点和输入节点彻底断开。
错误2:backward返回的梯度无需设置requires_grad=True
backward方法的返回值是梯度数值,属于计算图的末端结果,不是可微分节点,设置requires_grad=True完全多余,还会干扰梯度传播的正确性。同时要保证返回的梯度tensor与输入vertices的设备(CPU/GPU)、数据类型完全匹配,避免不兼容问题。
错误3:遗漏链式法则的上游梯度相乘
backward方法接收的grad_out是上游节点传来的梯度,必须与当前计算的grad_vertices相乘,这是链式法则的核心要求,之前的代码遗漏了这一步,会导致梯度计算结果完全错误。
修正后的代码
class Projector(torch.autograd.Function): @staticmethod def forward(ctx, vertices, non_torch_var1, non_torch_var2, non_torch_var3): ctx.save_for_backward(vertices) # 将非tensor变量打包存入ctx,更简洁规范 ctx.non_torch_vars = (non_torch_var1, non_torch_var2, non_torch_var3) vertices_np = vertices.detach().cpu().numpy() out_np = project_mesh(vertices_np, non_torch_var1, non_torch_var2, non_torch_var3) # 生成与输入匹配设备、类型的tensor,不手动设置requires_grad out = torch.tensor(out_np, device=vertices.device, dtype=vertices.dtype) return out @staticmethod def backward(ctx, grad_out): vertices, = ctx.saved_tensors non_torch_var1, non_torch_var2, non_torch_var3 = ctx.non_torch_vars vertices_np = vertices.detach().cpu().numpy() grad_vertices_np = grad_project_mesh(vertices_np, non_torch_var1, non_torch_var2, non_torch_var3) # 生成匹配输入的梯度tensor,不设置requires_grad,同时乘以上游梯度 grad_vertices = torch.tensor(grad_vertices_np, device=vertices.device, dtype=vertices.dtype) * grad_out # 与输入参数顺序对应,非可微分参数返回None return grad_vertices, None, None, None
内容的提问来源于stack exchange,提问作者Domenico
相关产品推荐
相关产品推荐

