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

如何在PyTorch中为多输入函数手动设置偏导数(仅对极点求导)

解决PyTorch自定义多输入Autograd函数的问题

先修正你的双输入乘法示例

你的代码核心错误是没有通过apply方法调用自定义的Autograd函数,直接实例化类只会得到对象,不会执行前向传播。修正后的代码如下:

import torch

class double_in(torch.autograd.Function):
    @staticmethod
    def forward(ctx, input, constant):
        ctx.save_for_backward(input, constant)
        output = input * constant
        return output
    
    @staticmethod
    def backward(ctx, grad_output):
        input, constant = ctx.saved_tensors
        # 对input的梯度是constant乘以grad_output,对constant不需要求导返回None
        return grad_output * constant, None

x = torch.rand(1, requires_grad=True)
# 必须用.apply()调用
out = double_in.apply(x, 5)
print("x = ", x, " out = ", out)

# 测试反向传播
out.backward()
print("x的梯度: ", x.grad)  # 应该输出5,符合预期

扩展到你的博士项目场景

针对你的需求:函数接收极点、初始条件、哈密顿矩阵三个输入,仅需对极点求偏导,其他输入视为常数。自定义Autograd函数的写法如下:

核心要点

  • 前向传播:保存反向传播中需要用到的张量(比如极点、初始条件、哈密顿矩阵,或仅保存计算梯度所需的变量)
  • 反向传播:返回对应输入的梯度——仅对极点返回推导好的解析梯度,另外两个输入返回None(表示不需要求导)

示例框架(适配你的场景)

import torch

class PoleEstimationFunction(torch.autograd.Function):
    @staticmethod
    def forward(ctx, poles, init_cond, hamiltonian):
        # 保存反向传播需要用到的所有张量
        ctx.save_for_backward(poles, init_cond, hamiltonian)
        
        # 替换成你的估计函数计算逻辑:输入三个参数,输出损失相关的估计值
        # 示例占位逻辑,实际替换成你的代码
        estimation = torch.sum(poles @ hamiltonian @ init_cond)
        return estimation
    
    @staticmethod
    def backward(ctx, grad_output):
        poles, init_cond, hamiltonian = ctx.saved_tensors
        
        # 替换成你推导的对poles的解析偏导数公式
        # 示例:假设对poles的梯度是 hamiltonian @ init_cond,再乘以grad_output的链式传播
        poles_grad = grad_output * (hamiltonian @ init_cond)
        
        # 初始条件和哈密顿矩阵不需要求导,返回None
        return poles_grad, None, None

# 使用示例
poles = torch.randn(3, requires_grad=True)  # 你的极点张量
init_cond = torch.randn(3)  # 初始条件,不需要求导
hamiltonian = torch.randn(3,3)  # 哈密顿矩阵,不需要求导

# 调用自定义函数
estimation = PoleEstimationFunction.apply(poles, init_cond, hamiltonian)

# 反向传播计算梯度
estimation.backward()

# 查看极点的梯度
print("极点的梯度: ", poles.grad)

关键注意事项

  1. backward方法的返回值数量必须和forward的输入参数数量一致,不需要求导的参数对应返回None
  2. 所有张量操作要保证维度匹配,避免广播错误
  3. 解析梯度的推导必须准确——这是自定义Autograd函数的核心,推导错误会导致反向传播梯度失效,直接影响模型训练

内容的提问来源于stack exchange,提问作者shmirrkk

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 00:43:21