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

自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 11:45:26