【问题标题】:Is there a `zero_grad()` function in Flux.jlFlux.jl 中是否有 `zero_grad()` 函数
【发布时间】:2021-06-26 18:44:53
【问题描述】:

在 PyTorch 中,您通常必须在进行反向传播之前将梯度归零。在 Flux 中是这种情况吗?如果是这样,这样做的程序化方式是什么?

【问题讨论】:

    标签: julia flux.jl


    【解决方案1】:

    tl;博士

    不,没有必要。

    解释

    Flux 用于使用 Tracker,这是一种微分系统,其中每个被跟踪的数组都可以保存一个梯度。我认为这是与 pytorch 类似的设计。反向传播两次可能会导致归零旨在避免的问题(尽管默认值试图保护您):

    julia> using Tracker
    
    julia> x_tr = Tracker.param([1 2 3])
    Tracked 1×3 Matrix{Float64}:
     1.0  2.0  3.0
    
    julia> y_tr = sum(abs2, x_tr)
    14.0 (tracked)
    
    julia> Tracker.back!(y_tr, 1; once=false)
    
    julia> x_tr.grad
    1×3 Matrix{Float64}:
     2.0  4.0  6.0
    
    julia> Tracker.back!(y_tr, 1; once=false) # by default (i.e. with once=true) this would be an error
    
    julia> x_tr.grad
    1×3 Matrix{Float64}:
     4.0  8.0  12.0
    

    现在它使用 Zygote,它不使用跟踪数组类型。相反,要跟踪的评估必须通过调用Zygote.gradient 进行,然后它可以查看和操作源代码以编写新的渐变代码。对它的重复调用每次都会产生相同的梯度;没有需要清理的存储状态。

    julia> using Zygote
    
    julia> x = [1 2 3]  # an ordinary Array
    1×3 Matrix{Int64}:
     1  2  3
    
    julia> Zygote.gradient(x -> sum(abs2, x), x)
    ([2 4 6],)
    
    julia> Zygote.gradient(x -> sum(abs2, x), x)
    ([2 4 6],)
    
    julia> y, bk = Zygote.pullback(x -> sum(abs2, x), x);
    
    julia> bk(1.0)
    ([2.0 4.0 6.0],)
    
    julia> bk(1.0)
    ([2.0 4.0 6.0],)
    

    Tracker也可以这样使用,而不是自己处理paramback!

    julia> Tracker.gradient(x -> sum(abs2, x), [1, 2, 3])
    ([2.0, 4.0, 6.0] (tracked),)
    

    【讨论】:

      猜你喜欢
      • 2020-12-06
      • 1970-01-01
      • 1970-01-01
      • 2021-11-06
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2018-03-13
      相关资源
      最近更新 更多