Fix dβ in the Enzyme rule for tensoradd! when an Active β is zero - #307
Conversation
…stead of its activity
| # form caches if needed | ||
| cache_A = EnzymeRules.overwritten(config)[3] ? copy(A_dA.val) : nothing | ||
| cache_C = !iszero(β_dβ.val) ? copy(C_dC.val) : C_dC.val | ||
| cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : nothing |
There was a problem hiding this comment.
Could we keep the C_dC.val here for the second option for type stability?
There was a problem hiding this comment.
Sure!
Just to be sure: I copied this exactly from the tensortrace! rule, since I'm not entirely comfortable with what I'm doing here it seemed safest to directly copy an existing pattern. Would it then make sense to also change the second option in the tensortrace! rule?
There was a problem hiding this comment.
I'm just trying to make sure the type of cache_C doesn't change in this rule, probably we can update the trace one separately.
There was a problem hiding this comment.
Is type stability really an issue here? isa(β_dβ, Const) is a check in type domain, so the two cases of the ? : , while leading to different types, should also be associated with different types of the input arguments, i.e. still type stable?
There was a problem hiding this comment.
Sorry, I've phrased this poorly, you're right. I think it makes things a bit easier on the compiler when cache can be proven to only ever have the same type as C.val (in this case), not C.val or nothing. Then we should only ever need to compile versions of the EnzymeRules.reverse for cache having that type. Maybe this isn't as much of a problem now that things in Enzyme internals on 1.12+ are improving (so we don't compile the entire world twice)? But on 1.10 it might matter. I also find the version with copy(C.val) : C.val a little easier to read, but others might have a different opinion :)
Jutho
left a comment
There was a problem hiding this comment.
Approve but see my comment. I am fine with either choice for the false case, it shouldn't matter, but prefer consistency with the other methods.
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
|
ROCm queue is very backed up (16h) and will not be affected by this so I will just merge |
I somehow stumbled into (what I think is) a bug in an Enzyme rule while messing around with implicit differentiation.
Bug
The Enzyme reverse rule of
tensoradd!decides whether to cacheCfromβ's value:Cis only needed fordβ = ⟨C, ΔC⟩, which must useC's value before the call. With anActiveβthat happens to be0.0,Cis not cached, anddβis computed from theoverwritten
C: a silently wrong gradient. In the other direction, aConstnonzeroβ(for exampleOne()when accumulating into an existing tensor) copiesCfor nothing.The
tensorcontract!andtensortrace!rules already useβ's activity instead:!isa(β_dβ, Const).Reproducer
f(C, β) = sum(abs2, tensoradd!(C, A, p, false, 1, β)), whose exact derivative isdf/dβ = 2⟨C₀, A + βC₀⟩.Reproducer
On
main(1250807):With this PR:
The existing tests never catch this: an
Activeβis always random, and a zeroβis alwaysZero(), which isConst.Fix
The same condition the other two rules use. When
βisConst, the cache is only passed totensoradd_pullback_dA!andtensoradd_pullback_dα, which don't use it.Tests
test/enzyme.jlgains atensoradd!case with(0.0, Active)and one with(randn(), Const). Onmainthe first fails (25/26: the reverse-mode gradient disagrees withfinite differences, 39.215 vs 10.314) and the second passes; with the fix both pass (52/52).