【问题标题】:How to return intermideate gradients (for non-leaf nodes) in pytorch?如何在pytorch中返回中间梯度(对于非叶节点)?
【发布时间】:2019-03-22 17:51:05
【问题描述】:

我的问题是关于 pytorch register_hook 的语法。

x = torch.tensor([1.], requires_grad=True)
y = x**2
z = 2*y

x.register_hook(print)
y.register_hook(print)

z.backward()

输出:

tensor([2.])
tensor([4.])

这个 sn-p 只是分别打印 z w.r.t xy 的梯度。

现在我(很可能是微不足道的)问题是如何返回中间渐变(而不仅仅是打印)?

更新:

看来调用retain_grad() 解决了叶节点的问题。前任。 y.retain_grad().

但是,retain_grad 似乎无法解决非叶节点的问题。有什么建议吗?

【问题讨论】:

    标签: python gradient pytorch


    【解决方案1】:

    我认为您可以使用这些钩子将渐变存储在全局变量中:

    grads = []
    x = torch.tensor([1.], requires_grad=True)
    y = x**2 + 1
    z = 2*y
    
    x.register_hook(lambda d:grads.append(d))
    y.register_hook(lambda d:grads.append(d))
    
    z.backward()
    

    但您很可能还需要记住计算这些梯度的相应张量。在这种情况下,我们使用dict 代替list 稍微扩展一下:

    grads = {}
    x = torch.tensor([1.,2.], requires_grad=True)
    y = x**2 + 1
    z = 2*y
    
    def store(grad,parent):
        print(grad,parent)
        grads[parent] = grad.clone()
    
    x.register_hook(lambda grad:store(grad,x))
    y.register_hook(lambda grad:store(grad,y))
    
    z.sum().backward()
    

    例如,现在您可以使用 grads[y] 访问张量 y 的 grad

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2015-08-14
      • 2013-02-15
      • 2020-06-10
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多