# Idea to make Zygote support mutation in easy cases

**URL:** <https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118>\
**Category:** Machine Learning\
**Tags:** zygote, ad\
**Created:** [December 26, 2022, 12:51am UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118 "2022-12-26T00:51:57Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![Lilith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lilith/32/27492_2.png) [@Lilith](https://discourse.julialang.org/u/Lilith)\
**Post date:** [December 26, 2022, 12:51am UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/1 "2022-12-26T00:51:57Z")

</div>

# Motivation

The appeal of automatic differentiation is that you can write a function using ordinary Julia and the gradient is automatically computed. Disallowing mutation in that function seems reasonable on the surface—it seems like it should be pretty easy to write code that doesn’t mutate arrays.

The problem is that mutation is tucked away in many places throughout the Julia ecosystem, in many cases as an implementation detail that is not visible to the user of that functionality. For example, `Float64[f(x) for x in input]` mutates. To quote the [Zygote documentation](https://fluxml.ai/Zygote.jl/dev/limitations/) “Non-mutating functions may also use mutation under the hood. This can be done for performance reasons or code re-use.”

# Idea

The fundamental issue with mutation is that you loose information when you overwrite a value. The idea is to support mutation when the value prior to the mutation is not used in gradient calculation. In my estimation, this would fix most of the hard to spot or unexpected mutations because Zygote would no longer throw for operations that mutate in a manner that is irrelevant to gradient calculation (e.g. use of mutation for code re-use within a library function). Further, in all or almost all cases it should be possible to add a `copy` operation just before the mutation to avoid the error if performance isn’t a major concern.

# Implementation

An array `a::T` where `T<:AbstractArray` could be represented as `mutable struct WrappedArray{T}; const a::T; valid::Bool; end` on the backwards pass with `valid` starting out as true. Pullback for mutating operations would set `valid` to false and access to `a` for gradient computation would be gated by runtime validity checking. The runtime cost of these checks would be present even if there was no mutation, but should be negligible arrays of more than a few elements. Ideally someone more familiar with Zygote.jl would be able to devise a system that has no runtime cost in most cases.

cc @MikeInnes who has [previously approached this issue](https://github.com/FluxML/Zygote.jl/pull/75)  
cc @ToucheSir, this proposal hopes to begin to address your comment [here](https://github.com/FluxML/Zygote.jl/issues/1343#issuecomment-1364726514)

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [December 26, 2022, 5:12am UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/2 "2022-12-26T05:12:32Z")

</div>

Have you checked Zygote’s buffers? It seems they do exactly what you suggest.

---

<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:** [December 26, 2022, 1:09pm UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/3 "2022-12-26T13:09:32Z")

</div>

Mutation is confusing but this sounds similar to a proposal in [ChainRules#521](https://github.com/JuliaDiff/ChainRules.jl/pull/521), where the idea is to make `return fill!(similar(x), y)` work by giving `fill!` a rule which poisons the gradient of its first argument.

But Zygote (1) at present doesn’t call the pullback at all when the function’s return is not used (as is common in mutating paths), and (2) since that issue seems to have been taught to ignore ChainRules’s not-implemented mechanism.

It’s possible that this has other problems too.

---

<div class="post-metadata">

**Author:** ![Lilith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lilith/32/27492_2.png) [@Lilith](https://discourse.julialang.org/u/Lilith)\
**Post date:** [December 26, 2022, 3:52pm UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/4 "2022-12-26T15:52:15Z")

</div>

> [@Tomas\_Pevny](#):
>
> Have you checked Zygote’s buffers? It seems they do exactly what you suggest.

I want to differentiate generic Julia code that predates or is otherwise unconcerned with compatibility with Zygote. For example, I can’t expect StatsBase’s `mean` function to use Zygote’s buffers as a way to fix [gradient() fails on array mutation for `mean(f, x; dims)` · Issue #1128 · FluxML/Zygote.jl · GitHub](https://github.com/FluxML/Zygote.jl/issues/1128).

> [@mcabbott](#):
>
> make `return fill!(similar(x), y)` work by giving `fill!` a rule which poisons the gradient of its first argument.

There is a small but pivotal difference between that proposal and my own. Rather than poison the gradient, I poison the data and set the gradient to zero (ideally a structural zero). A key objection to [ChainRules#521](https://github.com/JuliaDiff/ChainRules.jl/pull/521) is “that it will cause any other rule which has captured `x` to give wrong answers” (@mcabbott [here](https://github.com/JuliaDiff/ChainRules.jl/pull/521#issuecomment-909864618)). By poisoning the data, any other rule which captures the mutated value prior to mutation will throw.

---

<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:** [December 26, 2022, 4:44pm UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/5 "2022-12-26T16:44:04Z")

</div>

Yes, after thinking some more, passing something back in the gradient won’t always be enough. That’s what 521 looked at, and what I read this as saying, alter the gradient representation:

> [@Lilith](#):
>
> An array `a::T` […] could be represented as `mutable struct WrappedArray{T}; const a::T; valid::Bool; end` on the backwards pass

Now you seem to be saying that the object with this extra flag is present on the forward pass. In which case you can simulate the effect by having the pullback write NaN into the original `a`. (Or just restore the original values before mutation.)

What this still won’t solve is that, when the return value of `fill!(xs, y)` is discarded, its pullback won’t get the gradient for the new `xs`. In the spirit of trying to make simple cases work you could, modulo (1), (2) above, have the pullback return NotImplemented for both `dx` and `dy`. That could perhaps let something like `x[1]=0` work, when you don’t want the gradient of `x`. But won’t help for `sum(Float32[x for _ in 1:3])`.

---

<div class="post-metadata">

**Author:** ![Lilith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lilith/32/27492_2.png) [@Lilith](https://discourse.julialang.org/u/Lilith)\
**Post date:** [December 27, 2022, 11:00am UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/6 "2022-12-27T11:00:25Z")

</div>

> [@mcabbott](#):
>
> What this still won’t solve is that, when the return value of `fill!(xs, y)` is discarded, its pullback won’t get the gradient for the new `xs`.

I believe this is an example of what you are referring to, and yes, with my original proposal it would return `[0, 0, 0]`

```julia
gradient([1, 2, 3]) do x
    y = x*x
    fill!(x, zero(x))
    y
end

```

I don’t fully understand your point (2) above, but for point (1), the standard Julia compiler only elides functions whose return values are not used if they are free of side effects, could Zygote do a similar thing in pullbacks? That is, only elide pullbacks of functions whose return values are ignored when the pullback has no side effects. The trivial pullback of functions with `@nograd` is free of side effects.

> [@mcabbott](#):
>
> Now you seem to be saying that the object with this extra flag is present on the forward pass. In which case you can simulate the effect by having the pullback write NaN into the original `a`. (Or just restore the original values before mutation.)

Yes, this is closer to what I am suggesting. Unfortunately, writing `NaN` is not a viable option because not all arrays support `NaN` and restoring the original values is not ideal because that would require all mutating operations to make a copy. I don’t want `sum(Float32[x for _ in 1:3])` to allocate twice: first for the array that is populated and summed and then a copy of that uninitialized array to restore later. Ideally there would be some way to poison these arrays at compile time.

---

<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:** [December 27, 2022, 4:34pm UTC](https://discourse.julialang.org/t/idea-to-make-zygote-support-mutation-in-easy-cases/92118/7 "2022-12-27T16:34:21Z")

</div>

> [@Lilith](#):
>
> could Zygote do a similar thing in pullbacks? That is, only elide pullbacks of functions whose return values are ignored when the pullback has no side effects. The trivial pullback of functions with `@nograd` is free of side effects.

Zygote only sees untyped, pre-inference IR (think `@code_lowered` or `@code_warntype` without type information). It simply does not have enough information to do anything beyond the simplest, surface level transformations.

That said, it looks like we can force Zygote to evaluate pullbacks for mutating functions even if the result isn’t used. See [rrule for fill! by CarloLucibello · Pull Request #521 · JuliaDiff/ChainRules.jl · GitHub](https://github.com/JuliaDiff/ChainRules.jl/pull/521#issuecomment-1365304188).

> [@Lilith](#):
>
> Ideally there would be some way to poison these arrays at compile time.

One issue with doing this via a wrapper is that there are hundreds of existing ChainRules rrules (spread over numerous packages) out there which take array arguments. Making all of those aware of the wrapper type and able to check for validity/poison would be a massive effort, hence the exploration of less distruptive alternatives like writing NaNs. It’s quite possible this would require major changes to ChainRulesCore’s interface, as discussed in [Ability to specify different rules based on what combinations of inputs are actually being used · Issue #452 · JuliaDiff/ChainRulesCore.jl · GitHub](https://github.com/JuliaDiff/ChainRulesCore.jl/issues/452), [mutating calls · Issue #242 · JuliaDiff/ChainRulesCore.jl · GitHub](https://github.com/JuliaDiff/ChainRulesCore.jl/issues/242), etc.
