# Suggestions to improve Zygote performance for simple vector map/broadcast/comprehension?

**URL:** <https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207>\
**Category:** General Usage\
**Tags:** zygote\
**Created:** [February 27, 2020, 1:30am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207 "2020-02-27T01:30:38Z")\
**Posts on this page:** 7\
**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:** [February 27, 2020, 1:30am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/1 "2020-02-27T01:30:38Z")

</div>

I’ve got some code which is building some moderate sized vectors, which I’d like to derive through. Here’s an example building these vectors with a comprehension (the issue is the same if I use `map` or `broadcast`):

```julia
build_vector(x) = [i<500 ? x : 0 for i=1:1000]
@btime gradient(x -> sum(build_vector(x)), 1) # ~2ms

```

This is on a perfromance critical inner loop, and it turns out this is ~1000 times slower than if I wrote the adjoint by hand,

```julia
build_vector_with_adjoint(x) = build_vector(x)
@adjoint function build_vector_with_adjoint(x)
    y = build_vector(x)
    function back(Δ)
        b = [i<500 ? 1 : 0 for i=1:1000]
        (b'Δ,)
    end
    y, back
end
@btime gradient(x->sum(build_vector_with_adjoint(x)), 1) # ~2μs

```

I don’t think I’m cheating _too_ bad with this custom adjoint, it seems like this should basically be what Zygote should be writing for me. Profiling does show me some dynamic dispatch deep in the Zygote call-tree but I’m not familiar enough with the internals to make sense of it. The Zygote broadcast.jl source code has some comments alluding to performance hits and generic fallbacks, maybe I’m inadvertantly hitting something here? Any other suggestions to gain some performance without writing custom adjoints (which in my real non-MWE I think would be far more painful than here)? Thanks.

---

<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:** [February 27, 2020, 1:53am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/2 "2020-02-27T01:53:47Z")

</div>

Here are a few ideas, depending on how close this is to your non-MWE:

```julia
julia> build_vector(x) = [i<500 ? x : zero(x) for i=1:1000];

julia> @btime Zygote.gradient(x -> sum(build_vector(x)), 1)
  2.490 ms (15573 allocations: 644.61 KiB)
(499,)

julia> @btime ForwardDiff.derivative(x -> sum(build_vector(x)), 1)
  1.470 μs (1 allocation: 15.75 KiB)
499

julia> @btime Zygote.gradient(x -> sum(Zygote.forwarddiff(build_vector,x)), 1)
  6.590 μs (28 allocations: 40.34 KiB)
(499,)

```

---

<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:** [February 27, 2020, 2:10am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/3 "2020-02-27T02:10:05Z")

</div>

Thanks, thats helpful to see. In my non-MWE `x` is ~10 dimensional so I wanted to use reverse-mode, but maybe this still wins out, I can try it.

---

<div class="post-metadata">

**Author:** ![mhar](https://avatars.discourse-cdn.com/v4/letter/m/74df32/32.png) [@mhar](https://discourse.julialang.org/u/mhar)\
**Post date:** [November 8, 2022, 4:34pm UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/4 "2022-11-08T16:34:26Z")

</div>

Is there a way to re-write the original MWE (a replacement/substitute for list comprehension) so that reverse-mode is still fast? I have a similar problem with the same MWE but a much more complicated list comprehension that has many parameters.

---

<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:** [November 9, 2022, 1:47am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/5 "2022-11-09T01:47:56Z")

</div>

It depends on what code is running inside the comprehension, but probably. The key performance sink pitfall in the OP is that it uses control flow (conditionals and loops). Zygote isn’t able to generate efficient code for functions using control flow, so you’ll see both slower speeds and more allocations. Array comprehensions/map/broadcast with these functions is a worst-case scenario because it literally multiplies the overhead over the number of elements processed.

We can show the impact of removing control flow by using a branchless conditional (`ifelse`) instead of the ternary:

```julia
build_vector2(x) = [ifelse(i<500, x, 0) for i=1:1000]

julia> @btime gradient(x -> sum(build_vector(x)), 1);
  830.159 μs (6559 allocations: 285.09 KiB)

julia> @btime gradient(x -> sum(build_vector2(x)), 1);
  17.263 μs (44 allocations: 119.25 KiB)

```

However, some functions must use control flow. In that case, you have a few options:

1. Use [API · ChainRules](https://juliadiff.org/ChainRulesCore.jl/stable/api.html#Ignoring-gradients) around functions/code blocks that use control flow but don’t need to be differentiated.
2. Define your own `rrule`(s) for functions that use control flow. The advice in [Writing good rules · ChainRules](https://juliadiff.org/ChainRulesCore.jl/stable/rule_author/writing_good_rules.html) applies as always, but one additional concern here is to make sure the type of the returned _pullback function_ is stable. If it isn’t, you’ll run into many of the same issues as Zygote does.

---

<div class="post-metadata">

**Author:** ![mhar](https://avatars.discourse-cdn.com/v4/letter/m/74df32/32.png) [@mhar](https://discourse.julialang.org/u/mhar)\
**Post date:** [November 9, 2022, 2:28am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/6 "2022-11-09T02:28:31Z")

</div>

Is there a good reference for code that Zygote _can_ generate efficient code for?

---

<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:** [November 9, 2022, 3:13am UTC](https://discourse.julialang.org/t/suggestions-to-improve-zygote-performance-for-simple-vector-map-broadcast-comprehension/35207/7 "2022-11-09T03:13:54Z")

</div>

That’s difficult to quantify because it’s mostly a subtractive thing (if you do X, things will be slower). Aside from things that just [aren’t supported](https://fluxml.ai/Zygote.jl/latest/limitations/), the biggest performance pitfall I’ve seen outside of control flow is repeated indexing/`view` of a small section/single element of a large array. That allocates O(length(array)) on the backwards pass. That’s less of a codegen and more of a runtime perf issue, however. Accessing and setting properties on mutable structs within sight of Zygote will also be type unstable, though there you only pay for the cost of the dynamic dispatch (whereas generated pullbacks for control flow can allocate quite a bit more besides).
