PyTorch 0.3.1中autograd求导报错:differentiated input is unreachable
解决PyTorch 0.3.1中拼接变量的梯度获取问题
嘿,我明白你的困扰了——在PyTorch 0.3.1里,因为w只是x、y、z的拼接产物,但完全没参与f的计算,计算图里没建立起w到f的连接,所以直接求w的梯度才会报RuntimeError: differentiated input is unreachable的错误。不过要得到类似[1;1;1]的批量梯度结果,有两种简单的一行代码实现方式:
方法1:直接拼接已有梯度
既然你已经能正确拿到x、y、z的梯度(都是1),只需要把它们直接拼接起来就行,完全不用修改原有的计算逻辑:
import torch from torch.autograd import Variable # 原代码不变 x = Variable(torch.tensor([1.0]), requires_grad=True) y = Variable(torch.tensor([2.0]), requires_grad=True) z = Variable(torch.tensor([3.0]), requires_grad=True) w = torch.cat([x, y, z]) f = x + y + z f.backward() # 一行代码得到目标梯度 w_grad = torch.cat([x.grad, y.grad, z.grad]) print(w_grad) # 输出: tensor([1., 1., 1.])
方法2:让w参与计算图(更符合语义)
如果你希望直接通过w.grad得到结果,可以修改f的计算方式,让w成为计算图的一部分——因为w是x、y、z的拼接,w.sum()的结果和x+y+z完全一致,这样反向传播时就能直接计算w的梯度:
import torch from torch.autograd import Variable x = Variable(torch.tensor([1.0]), requires_grad=True) y = Variable(torch.tensor([2.0]), requires_grad=True) z = Variable(torch.tensor([3.0]), requires_grad=True) w = torch.cat([x, y, z]) f = w.sum() # 用w的sum替代直接相加,让w进入计算图 f.backward() # 直接获取w的梯度 print(w.grad) # 输出: tensor([1., 1., 1.])
两种方法都能实现你的需求,方法2更贴合“求w的梯度”的语义,方法1则不需要改动原有的计算逻辑,看你更倾向哪种即可。
内容的提问来源于stack exchange,提问作者user650261
相关产品推荐
相关产品推荐

