# Speed up compilation with Zygote and PINN

**URL:** <https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652>\
**Category:** Machine Learning\
**Created:** [May 31, 2023, 7:43am UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652 "2023-05-31T07:43:16Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![Matthieu\_BARREAU](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/matthieu_barreau/32/50368_2.png) [@Matthieu\_BARREAU](https://discourse.julialang.org/u/Matthieu_BARREAU)\
**Post date:** [May 31, 2023, 7:43am UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652/1 "2023-05-31T07:43:16Z")

</div>

Hi!

I am relatively new to Julia and I am trying to move all my research projects to this new ecosystem instead of Python.

However, I have trouble when my project deals with Physics informed neural network. Since I am working on how to improve PINN, I need to keep control of what is happening and I cannot use any package. I am trying on my own but I am facing some compilation time that are excessively long. Here is a very simple example:

```julia
using Flux, Zygote

x_train = [2; 0;; 0; 1] * rand(Float32, 2, 10)
target = rand(Float32, 1, 10)

model = Chain(
    Dense(2 => 15, tanh),
    Dense(15 => 15, tanh),
    Dense(15 => 1)
)

function get_residual(f::Chain)
    df(u) = Zygote.gradient(u -> sum(f(u)), u)[1]
    #ddf(u) = Zygote.gradient(u -> sum(df(u)), u)[1]
    return u -> df(u)[1] .+ (1.0f0 .- 2.0f0 .* f(u)) .* df(u)[2] #.- 0.001f0 .* ddf(u)[2]
end

r = get_residual(model)
@info "Getting residuals"
timed = @timed r(x_train)
@show timed.time

function withgradient(model, x_train)
    loss, grads = Flux.withgradient(model) do m
        y_hat = m(x_train)
        loss = Flux.Losses.mse(y_hat, target)
        penalty = Flux.Losses.mse(get_residual(m)(x_train), zeros(Float32, 1, 10))
        loss + penalty
    end
    (loss, grads)
end

@info "Getting gradients"
timed = @timed withgradient(model, x_train)
@show timed.time

```

It takes a very long time to compile first time, even if the model is very simple (25s for getting the residuals and 264s for the gradients). Moreover, if I uncomment the lines in the function `get_residual` then it almost never finishes compilation. Such a code was running in a couple of seconds in Tensorflow before.

Do you have some suggestions to improve this code? Is there anything I am doing wrong? I tried using Yota (which is faster in my case) but it cannot compute the gradients in the end…

Best,  
Matthieu

---

<div class="post-metadata">

**Author:** ![Matthieu\_BARREAU](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/matthieu_barreau/32/50368_2.png) [@Matthieu\_BARREAU](https://discourse.julialang.org/u/Matthieu_BARREAU)\
**Post date:** [May 31, 2023, 3:03pm UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652/2 "2023-05-31T15:03:12Z")

</div>

I found that using ForwardDiff for getting the residuals is slightly more efficient (but can deal with the additional viscosity).  
Any idea on how to use forward mode to get the residual and reverse mode for the grads?  
I saw this post: [How to achieve good performance with Zygote.pushforward on a neural network](https://discourse.julialang.org/t/how-to-achieve-good-performance-with-zygote-pushforward-on-a-neural-network/55971) which might solve this issue, however, it is outdated and the codes there are not working anymore.

---

<div class="post-metadata">

**Author:** ![de-souza](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/de-souza/32/43417_2.png) [@de-souza](https://discourse.julialang.org/u/de-souza)\
**Post date:** [May 31, 2023, 3:20pm UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652/3 "2023-05-31T15:20:06Z")

</div>

Hi,

Based on what was said in a recent topic, I would recommend the approach taken by [TaylorDiff.jl](https://github.com/JuliaDiff/TaylorDiff.jl) for higher order automatic differentiation. See the discussion and links in that topic:

> [@Is it possible to do Nested AD ~elegantly~ in Julia? (PINNs)](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888):
>
> To differentiate a loss function defined in terms of a network’s derivatives seems to have been an issue since forever, specially when Zygote is involved; see [[1]](https://discourse.julialang.org/t/how-to-use-gradient-of-neural-network-as-the-loss-function/50569/11) [[2]](https://discourse.julialang.org/t/current-status-of-nested-ad/70753) [[3]](https://discourse.julialang.org/t/flux-pinn-1d-burgers/93262) and many more over at Zygote’s git. In many of these threads (specially pre 2022) it is said that [the release of Diffractor.jl would correct this issue](https://discourse.julialang.org/t/gradient-calculation-in-pinn/61525/3). Now that it seems that [Diffractor is mostly dead](https://discourse.julialang.org/t/state-of-diffractor-jl/92959), what is left? In my case, because I am working with exotic architectures, it is not possible to use @ChrisRackauckas’s Neura…

---

<div class="post-metadata">

**Author:** ![Matthieu\_BARREAU](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/matthieu_barreau/32/50368_2.png) [@Matthieu\_BARREAU](https://discourse.julialang.org/u/Matthieu_BARREAU)\
**Post date:** [June 1, 2023, 7:57am UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652/4 "2023-06-01T07:57:19Z")

</div>

Thanks @de-souza for this reference. Based on this work, I propose the following (not really optimized) solution that seems to work well and fast. However, when we compare with the fully Zygote solution, there is a difference in the gradients that is not that negligible. But this is difficult to say which one is the best since we do not have any reference.

```julia
using Flux, Zygote, ForwardDiff, TaylorDiff, SliceMap

x_train = [2; 0;; 0; 1] * rand(Float32, 2, 10)
target = rand(Float32, 1, 10)

model = Chain(
    Dense(2 => 15, tanh),
    Dense(15 => 15, tanh),
    Dense(15 => 1)
)

function get_residual(f::Chain)
    df(u) = Zygote.gradient(u -> sum(f(u)), u)[1]
    return u -> df(u)[[1], :] .+ (1.0f0 .- 2.0f0 .* f(u)) .* df(u)[[2], :] # Cannot do higher order
end

function get_residual_forward(f::Chain)
    df(u) = ForwardDiff.gradient(u -> sum(f(u)), u)
    #ddf(u) = ForwardDiff.gradient(u -> sum(df(u)), u)
    return u -> df(u)[[1], :] .+ (1.0f0 .- 2.0f0 .* f(u)) .* df(u)[[2], :] #.- 0.001f0 .* ddf(u)[2]
end

function get_residual_Taylor(f::Chain, x)
    x = convert(Vector{Float32}, x)
    dfdt(u) = TaylorDiff.derivative(u -> sum(f(u)), u, [1.0f0, 0.0f0], 1)
    dfdx(u) = TaylorDiff.derivative(u -> sum(f(u)), u, [0.0f0, 1.0f0], 1)
    #dfdxx(u) = TaylorDiff.derivative(u -> sum(f(u)), u, [0.0f0, 1.0f0], 2)
    return dfdt(x) .+ (1.0f0 .- 2.0f0 .* f(x)) .* dfdx(x) #.- 0.001f0 .* dfdxx(x)
end

function get_residual_fd(f::Chain)
    ε = cbrt(eps(Float32))
    ε₁ = [ε; 0]
    ε₂ = [0; ε]
    V(x) = (1.0f0 .- 2.0f0 .* f(x))
    return x -> (f(x .+ ε₁) - f(x)) / ε .+ V(x) .* (f(x .+ ε₂) - f(x)) / ε
end

r = get_residual(model)
@info "Getting residuals (Zygote)"
timed = @timed r(x_train)
@show timed

r_forward = get_residual_forward(model)
@info "Getting residuals (ForwardDiff)"
timed = @timed r_forward(x_train)
@show timed

r_taylor(x) = get_residual_Taylor(model, x)
@info "Getting residuals (TaylorDiff)"
timed = @timed mapcols(x -> get_residual_Taylor(model, x), x_train)
@show timed

r_fd = get_residual_fd(model)
@info "Getting residuals (FiniteDifference)"
timed = @timed r_fd(x_train)
@show timed

@info "Getting gradients (Zygote)"
function withgradient_Zygote(model, x_train)
    loss, grads = Zygote.withgradient(model) do m
        y_hat = m(x_train)
        loss = Flux.Losses.mse(y_hat, target)
        penalty = Flux.Losses.mse(get_residual(m)(x_train), zeros(Float32, 1, 10))
        loss + penalty
    end
    (loss, grads)
end

timed_Zygote = @timed withgradient_Zygote(model, x_train)
@show timed_Zygote

@info "Getting gradients (TaylorDiff)"
function withgradient_taylor(model, x_train)
    loss, grads = Zygote.withgradient(model) do m
        y_hat = m(x_train)
        loss = Flux.Losses.mse(y_hat, target)
        r = mapcols(x -> get_residual_Taylor(model, x), x_train)
        penalty = Flux.Losses.mse(r, zeros(Float32, 1, 10))
        loss + penalty
    end
    (loss, grads)
end

timed_taylor = @timed withgradient_taylor(model, x_train)
@show timed_taylor

```

The solution with Zygote takes around 175sec to compile and run (first time) while the solution with TaylorDiff is about 10 times less (14sec).

One solution would be to compare on a simple example with an analytical solution but I am not sure I have the will for doing that now.

I hope that can help you too @de-souza , tell me what you think!

---

<div class="post-metadata">

**Author:** ![de-souza](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/de-souza/32/43417_2.png) [@de-souza](https://discourse.julialang.org/u/de-souza)\
**Post date:** [June 1, 2023, 1:01pm UTC](https://discourse.julialang.org/t/speed-up-compilation-with-zygote-and-pinn/99652/5 "2023-06-01T13:01:23Z")

</div>

Thanks a lot for your solution. Perhaps the gradients are different because the residuals are not exactly equal when computed with Taylor or Zygote. Anyways I will save your solution for when I need it so thanks again!
