PyTorch中如何让梯度通过min函数传递?
解决PyTorch中梯度无法传递到min结果的问题
原代码的问题在于:直接用torch.tensor([a, b, c])构造新张量时,会提取带梯度的a的数值,导致新张量脱离原计算图,最终d无法传递梯度给a。
修改方案如下:
- 将整数
b、c转换为与a同数据类型、同设备的PyTorch张量(无需设置requires_grad=True,因为它们不需要梯度) - 使用
torch.stack(或torch.cat)将a、b、c的张量组合,确保a的梯度信息被保留在计算图中 - 再对组合后的张量计算最小值
修改后的代码:
import torch a = torch.tensor([4.], requires_grad=True) # 将b、c转为与a匹配的tensor b = torch.tensor(5., dtype=a.dtype, device=a.device) c = torch.tensor(6., dtype=a.dtype, device=a.device) # 用stack组合张量,保留梯度信息 d = torch.min(torch.stack([a, b, c])) # 验证梯度传递 d.backward() print(a.grad) # 输出tensor([1.]),说明梯度已成功传递
如果a是标量张量(如示例中的[4.]),也可以简化组合方式:
d = torch.min(torch.tensor([a, b, c], dtype=a.dtype, device=a.device))
这种方式同样能保留a的梯度,因为构造张量时指定了与a一致的类型和设备,且a作为张量元素被直接包含。
内容的提问来源于stack exchange,提问作者learner
相关产品推荐
相关产品推荐

