自定义autograd.Function时,LibTorch对应PyTorch ctx.needs_input_grad的API是什么?
问题描述
将自定义PyTorch autograd.Function 的Python代码转换为LibTorch C++代码时,发现LibTorch的AutogradContext没有对应PyTorch中ctx.needs_input_grad的属性/方法。尝试用v_dirs.required_grad()替代,但不确定是否正确,想知道正确的对应API是什么。
对应的Python代码如下:
class _SphericalHarmonics(torch.autograd.Function): """Spherical Harmonics""" @staticmethod def forward( ctx, sh_degree: int, dirs: Tensor, coeffs: Tensor, masks: Tensor ) -> Tensor: colors = _make_lazy_cuda_func("compute_sh_fwd")(sh_degree, dirs, coeffs, masks) ctx.save_for_backward(dirs, coeffs, masks) ctx.sh_degree = sh_degree ctx.num_bases = coeffs.shape[-2] return colors @staticmethod def backward(ctx, v_colors: Tensor): dirs, coeffs, masks = ctx.saved_tensors sh_degree = ctx.sh_degree num_bases = ctx.num_bases compute_v_dirs = ctx.needs_input_grad[1] v_coeffs, v_dirs = _make_lazy_cuda_func("compute_sh_bwd")( num_bases, sh_degree, dirs, coeffs, masks, v_colors.contiguous(), compute_v_dirs, ) if not compute_v_dirs: v_dirs = None return None, v_dirs, v_coeffs, None
解答
在LibTorch中,对应Python里ctx.needs_input_grad[idx]的是torch::autograd::AutogradContext类的needs_input_grad(size_t idx)成员函数,传入输入参数的索引即可判断该输入是否需要计算梯度。
对应你代码中Python侧的compute_v_dirs = ctx.needs_input_grad[1],C++侧应写成:
bool compute_v_dirs = ctx.needs_input_grad(1);
你尝试的v_dirs.required_grad()是错误的——v_dirs是反向传播的输出梯度张量,它的required_grad()属性和输入参数是否需要梯度没有直接关联,不能用来替代needs_input_grad的作用。
额外注意:
- 索引规则和Python一致,从0开始计数,代码中索引1对应
dirs参数,和原Python逻辑完全匹配。 - 如果某个输入不需要梯度,反向传播时返回
torch::Tensor()即可(对应Python中的None)。
内容的提问来源于stack exchange,提问作者Zhiguo Tang
相关产品推荐
相关产品推荐

