# Zygote, Flux, and Interpolations

**URL:** <https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824>\
**Category:** Machine Learning\
**Created:** [May 25, 2021, 11:05pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824 "2021-05-25T23:05:19Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![kiranshila](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kiranshila/32/25475_2.png) [@kiranshila](https://discourse.julialang.org/u/kiranshila)\
**Post date:** [May 25, 2021, 11:05pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/1 "2021-05-25T23:05:19Z")

</div>

Hey everyone!

I’m working on radio telescope imaging and am trying to do image reconstruction with an iterative, maximum entropy method. Many papers suggest using conjugate gradient, and have intricate derivations of gradients. I thought it would be neat to do the same thing, but pushing the model through some sort of Autodiff.

In this process, I am needing to take samples of an interpolated FFT and calculate a mean squared error to measured radio telescope data.

Something like this

```julia
function vis_res(image, vis_data, uv)
    # Generate interpolated visibilities from the image
    N = length(vis_data)
    image_fft = fft(image) |> fftshift
    freqs = fftfreq(size(image)[1]) |> fftshift
    interpolation = LinearInterpolation((freqs, freqs), image_fft)
    vis_interp = [interpolation(uv[:,i]...) for i ∈ 1:N]
    # Calculate visibility residuals
    return (abs.(vis_interp .- vis_data)).^2
end

```

However, Flux/Zygote doesn’t seem to be happy with the indexing of the interpolation. Throwing the error:

```julia
ERROR: ArgumentError: unable to check bounds for indices of type Interpolations.WeightedAdjIndex{2, Float64}

```

Any help in this regard would be greatly appreciated.

---

<div class="post-metadata">

**Author:** ![sefffal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sefffal/32/23640_2.png) [@sefffal](https://discourse.julialang.org/u/sefffal)\
**Post date:** [June 1, 2021, 10:12pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/2 "2021-06-01T22:12:51Z")

</div>

Funny, I just ran into the same error. I’m trying to differentiate through a likelihood function that uses an interpolation using Zygote.

I expect the incompatibility is within the Interpolations library rather than the indexing you show. My likelihood function also calls a manual bi-linear interpolation function I wrote, and it has no issue with that as far as I can tell.

---

<div class="post-metadata">

**Author:** ![sefffal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sefffal/32/23640_2.png) [@sefffal](https://discourse.julialang.org/u/sefffal)\
**Post date:** [June 1, 2021, 10:14pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/3 "2021-06-01T22:14:30Z")

</div>

You might also want to try testing ForwardDiff unless you have to use Zygote for some reason.

---

<div class="post-metadata">

**Author:** ![kiranshila](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kiranshila/32/25475_2.png) [@kiranshila](https://discourse.julialang.org/u/kiranshila)\
**Post date:** [June 1, 2021, 10:20pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/4 "2021-06-01T22:20:05Z")

</div>

Yeah I dug in deep to this - and it actually does work on the current master branch of Interpolations, as the release hasn’t been bumped in a few months. There is still a strange indexing problem, stemming from strangeness in the dimensionality of the gradient. I solved my specific example by just providing an explicit gradient.

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [June 1, 2021, 10:20pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/5 "2021-06-01T22:20:25Z")

</div>

It would help to have a full stacktrace as well. As-is (having no knowledge of how interpolations works), I have no idea which line is even causing the error.

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [June 1, 2021, 10:23pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/6 "2021-06-01T22:23:22Z")

</div>

~~Looks like Interpolations didn’t gain [AD support](https://github.com/JuliaMath/Interpolations.jl/pull/414) until after the latest release, so that makes sense~~ Edit: saw you commented already on a linked issue :). Worth reporting an issue if you’re getting incorrect gradient values.

---

<div class="post-metadata">

**Author:** ![sefffal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sefffal/32/23640_2.png) [@sefffal](https://discourse.julialang.org/u/sefffal)\
**Post date:** [November 14, 2023, 9:54pm UTC](https://discourse.julialang.org/t/zygote-flux-and-interpolations/61824/7 "2023-11-14T21:54:53Z")

</div>

For those finding your way here from Google, the following function gives a bilinear interpolation over the axes of a matrix and is easily auto differentiable with Zygote and others. You should generally prefer the battle-tested Interpolations.jl implementations where you can, but this one works in a pinch.

```julia

function bilininterp(data::AbstractMatrix, x, y)
    x1 = floor(Int,x)
    x2 = ceil(Int,x)
    y1 = floor(Int,y)
    y2 = ceil(Int,y)

    # Handle case of perfect grid alignment
    if x1 == x2
        x2 = x1 + 1
    end
    if y1 == y2
        y2 = y1 + 1
    end

    # Handle boundary conditions
    if x1 == 0
        x1 += 1
        x2 += 1
    end
    if y1 == 0
        y1 += 1
        y2 += 1
    end
    if x1 == size(data,1)
        x1 -= 1
        x2 -= 1
    end
    if y1 == size(data,2)
        y1 -= 1
        y2 -= 1
    end

    mat = [
        1 x1 y1 (x1 * y1)
        1 x1 y2 (x1 * y2)
        1 x2 y1 (x2 * y1)
        1 x2 y2 (x2 * y2)
    ]
    # Get surrounding pixels if finite (fallback to image median otherwise)
    col = [
        data[x1, y1]
        data[x1, y2]
        data[x2, y1]
        data[x2, y2]
    ]
    # Solve the system of equations without inverting
    coeffs = mat \ col
    interp =
        coeffs[1] + coeffs[2] * x + coeffs[3] * y + coeffs[4] * x * y
    return interp
end

```
