# Help optimizing 1000x Zygote overhead for linear interpolation

**URL:** https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549
**Category:** General Usage
**Created:** [April 15, 2022, 7:53pm UTC](https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549 "2022-04-15T19:53:40Z")
**Posts on this page:** 4
**Page:** 1

<div class="post-metadata">

### Author: ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)
#### Post date: [April 15, 2022, 7:53pm UTC](https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549/1 "2022-04-15T19:53:40Z")

</div>

I need to compute a gradient of an interpolation w.r.t. to the knot y-positions. None of the standard packages seemed to work with Zygote so I coded this very simple MWE:

```julia
using Zygote, ForwardDiff, BenchmarkTools

function LinearInterpolation(xdat::AbstractVector, ydat::AbstractVector{T}, extrapolation_value::T = T(NaN)) where {T}
    m = diff(ydat) ./ diff(xdat)
    x_lower, x_upper = first(xdat), last(xdat)
    function (x)
        if x_lower <= x <= x_upper
            # sets i such that x is between xdat[i] and xdat[i+1]
            i = Zygote.@ignore(max(1, searchsortedfirst(xdat, x) - 1))
            @inbounds(ydat[i] + m[i]*(x-xdat[i]))
        else
            return extrapolation_value
        end
    end
end

```

However, while the Zygote gradient works, its incredibly slow, its about ~1000X slower than the evaluation itself, and ~100X slower than with ForwardDiff for a typical use-case for me:

```julia
xdat = collect(range(0,1,length=100))
ydat = rand(100)
x = rand(128,128)

@btime sum(LinearInterpolation($xdat, $ydat).($x))
# ~500μs

@btime Zygote.gradient(ydat -> sum(LinearInterpolation($xdat, ydat).($x)), $ydat)
# ~500ms

@btime ForwardDiff.gradient(ydat -> sum(LinearInterpolation($xdat, ydat).($x)), $ydat)
# ~5ms

```

Profiling reveals lots of time spent in Zygote internals which I’m not familiar with (eg `_generate_pullback_via_decomposition`, dynamic dispatch, and stuff where stack frames dissapear, so I guess are outside of Julia?) so that didn’t get me far.

I’m sure I’m just hitting a case where Zygote is known to be bad, but I’m wondering if anyone could still give some hints how I might rewrite this to make it less disastorously slow?

Alternatively, I guess the nuclear option is code the rrule for the entire function by hand. However, what’s the rrule of a function which returns a closure? I suppose I could also use ForwardDiff in the rrule, any examples of how to do that? Thanks.

---

<div class="post-metadata">

### Author: ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)
#### Post date: [April 15, 2022, 8:02pm UTC](https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549/2 "2022-04-15T20:02:09Z")

</div>

I recently learned that you can mix reverse and forward differentiation. For Zygote [this](https://docs.juliahub.com/Zygote/4kbLI/0.5.4/utils/#Zygote.forwarddiff) is the magic incantation.

---

<div class="post-metadata">

### Author: ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)
#### Post date: [April 15, 2022, 8:57pm UTC](https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549/3 "2022-04-15T20:57:47Z")

</div>

I’m not sure about this but have you tried simply removing the closure to see how this affects the running time?

---

<div class="post-metadata">

### Author: ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)
#### Post date: [April 15, 2022, 11:16pm UTC](https://discourse.julialang.org/t/help-optimizing-1000x-zygote-overhead-for-linear-interpolation/79549/4 "2022-04-15T23:16:35Z")

</div>

> [@gdalle](#):
>
> I’m not sure about this but have you tried simply removing the closure to see how this affects the running time?

Just checked, doesn’t seem to.

> [@goerch](#):
>
> I recently learned that you can mix reverse and forward differentiation. For Zygote [this](https://docs.juliahub.com/Zygote/4kbLI/0.5.4/utils/#Zygote.forwarddiff) is the magic incantation.

Didn’t know about that one, indeed that makes it super easy to quickly switch to ForwardDiff for the relevant part. That’s gets me on the order of the ForwardDiff result I had above, which may be OK for now. Thanks! Here’s what I have now for reference (combined into one function):

```julia
function LinearInterpolation(xdat::AbstractVector{TX}, ydat::AbstractVector{TY}, x::AbstractArray{TX}, extrapolation_value::TY = TY(NaN)) where {TX,TY}
    x_lower, x_upper = first(xdat), last(xdat)
    Zygote.forwarddiff(ydat) do ydat
        m = diff(ydat) ./ diff(xdat)
        map(x) do x
            if x_lower <= x <= x_upper
                # sets i such that x is between xdat[i] and xdat[i+1]
                i = max(1, searchsortedfirst(xdat, x) - 1)
                @inbounds(ydat[i] + m[i]*(x-xdat[i]))
            else
                extrapolation_value
            end
        end
    end
end

```
