# Adjoint for threaded map (ThreadsX.map)

**URL:** <https://discourse.julialang.org/t/adjoint-for-threaded-map-threadsx-map/97142>\
**Category:** Performance\
**Tags:** question, parallel, multithreading, zygote, threads\
**Created:** [April 5, 2023, 10:59pm UTC](https://discourse.julialang.org/t/adjoint-for-threaded-map-threadsx-map/97142 "2023-04-05T22:59:28Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![mfishelson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mfishelson/32/48794_2.png) [@mfishelson](https://discourse.julialang.org/u/mfishelson)\
**Post date:** [April 5, 2023, 10:59pm UTC](https://discourse.julialang.org/t/adjoint-for-threaded-map-threadsx-map/97142/1 "2023-04-05T22:59:28Z")

</div>

I am using Zygote to compute the gradient of a function which calls `map(f, x)`, where applying `f` to each element of the array `x` is slow. I would like to speed up the gradient computation by using a threaded map such as `ThreadsX.map`; however, Zygote cannot differentiate through this.

Thus, I would like to write a custom adjoint for `ThreadsX.map` that also parallelizes the backwards pass (using threads) but don’t know how to do this and would be grateful for any help.  
I found [this discussion](https://discourse.julialang.org/t/parallel-reductions-with-zygote/75969) to be helpful – it is essentially the same problem but for `ThreadsX.sum`. The solution implemented a custom adjoint for `ThreadsX.sum` by taking the regular rrule for sum and replacing the map and sum calls with ThreadsX versions. I tried to do something analogous to this but couldn’t find the rrule for map in `ChainRules.jl`.

Alternatively, is there another package that contains a threaded map that is compatible with Zygote?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [April 6, 2023, 2:02am UTC](https://discourse.julialang.org/t/adjoint-for-threaded-map-threadsx-map/97142/2 "2023-04-06T02:02:43Z")

</div>

> [@mfishelson](#):
>
> but couldn’t find the rrule for map in `ChainRules.jl`.

There is [this one, for `map(f, ::Tuple)`](https://github.com/JuliaDiff/ChainRules.jl/blob/77ef0eb15fdd207028c059927f2456819d62df8c/src/rulesets/Base/base.jl#L210-L245), and Zygote has [this](https://github.com/FluxML/Zygote.jl/blob/413728b65ff0d9a891afe2670ffac2d3a7ccf742/src/lib/array.jl#L166-L229). The basic idea is quite simple, but there are elaborations to deal with `map(f, x, y, z)`, and to be more efficient in some cases.

```julia
function ChainRulesCore.rrule(config::RuleConfig{>:HasReverseMode}, ::typeof(my_map), f, X::AbstractArray)
    hobbits = my_map(X) do x # this makes an array of tuples
        y, back = rrule_via_ad(config, f, x)
    end
    Y = map(first, hobbits)
    function map_pullback(dY_raw)
        dY = unthunk(dY_raw)
        # Should really do these in the reverse order
        backevals = my_map(hobbits, dY) do (y, back), dy
            dx, dx = back(dy)
        end
        df = ProjectTo(f)(sum(first, backevals))
        dX = map(last, backevals)
        return (NoTangent(), df, dX)
    end
    return Y, map_pullback
end

my_map(f, Xs...) = map(@show(f), Xs...)

gradient(x -> sum(map(inv, x)), [1,2,3.0])
gradient(x -> sum(my_map(inv, x)), [1,2,3.0]) # dX

gradient(x -> sum(map(z -> z/x, 1:3)), 4.0)
gradient(x -> sum(my_map(z -> z/x, 1:3)), 4.0) # df

```
