Skip to content

Fix dβ in the Enzyme rule for tensoradd! when an Active β is zero - #307

Merged
kshyatt merged 3 commits into
mainfrom
lb/enzyme-tensoradd-beta
Sep 25, 2026
Merged

kshyatt merged 3 commits into
mainfrom
lb/enzyme-tensoradd-beta

Conversation

@leburgel

Copy link
Copy Markdown
Member

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 cache C from β's value:

cache_C = !iszero(β_dβ.val) ? copy(C_dC.val) : C_dC.val

C is only needed for dβ = ⟨C, ΔC⟩, which must use C's value before the call. With an
Active β that happens to be 0.0, C is not cached, and dβ is computed from the
overwritten C: a silently wrong gradient. In the other direction, a Const nonzero
β (for example One() when accumulating into an existing tensor) copies C for nothing.

The tensorcontract! and tensortrace! rules already use β's activity instead:
!isa(β_dβ, Const).

Reproducer

f(C, β) = sum(abs2, tensoradd!(C, A, p, false, 1, β)), whose exact derivative is
df/dβ = 2⟨C₀, A + βC₀⟩.

Reproducer
using Enzyme, TensorOperations, LinearAlgebra, Random
Random.seed!(1)

A = randn(3, 4); C0 = randn(3, 4); Id = Matrix(1.0I, 4, 4)
# C′ = A + β C0 and f = |C′|², so df/dβ = 2⟨C0, C′⟩
add!(C, β) = tensoradd!(C, A, ((1, 2), ()), false, 1.0, β)
contract!(C, β) = tensorcontract!(C, A, ((1,), (2,)), false, Id, ((1,), (2,)), false, ((1,), (2,)), 1.0, β)
for (name, op!) in (("tensoradd!", add!), ("tensorcontract!", contract!)), β in (0.5, 0.0)
    f(C, β) = (op!(C, β); sum(abs2, C))
    dβ = Enzyme.autodiff(Reverse, f, Active, Duplicated(copy(C0), zero(C0)), Active(β))[1][2]
    exact = 2 * dot(C0, A .+ β .* C0)
    println(rpad(name, 16), "β = $β: dβ = $dβ, exact $exact", dβ ≈ exact ? "" : "   <-- wrong")
end

On main (1250807):

tensoradd!      β = 0.5: dβ = 16.431182701766023, exact 16.431182701766023
tensoradd!      β = 0.0: dβ = 17.55833018081556, exact 2.139868546270043   <-- wrong
tensorcontract! β = 0.5: dβ = 16.431182701766023, exact 16.431182701766023
tensorcontract! β = 0.0: dβ = 2.139868546270043, exact 2.139868546270043

With this PR:

tensoradd!      β = 0.5: dβ = 16.431182701766023, exact 16.431182701766023
tensoradd!      β = 0.0: dβ = 2.139868546270043, exact 2.139868546270043
tensorcontract! β = 0.5: dβ = 16.431182701766023, exact 16.431182701766023
tensorcontract! β = 0.0: dβ = 2.139868546270043, exact 2.139868546270043

The existing tests never catch this: an Active β is always random, and a zero β is always
Zero(), which is Const.

Fix

cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : nothing

The same condition the other two rules use. When β is Const, the cache is only passed to
tensoradd_pullback_dA! and tensoradd_pullback_dα, which don't use it.

Tests

test/enzyme.jl gains a tensoradd! case with (0.0, Active) and one with
(randn(), Const). On main the first fails (25/26: the reverse-mode gradient disagrees with
finite differences, 39.215 vs 10.314) and the second passes; with the fix both pass (52/52).

@leburgel
leburgel requested a review from kshyatt September 24, 2026 13:31
# 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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we keep the C_dC.val here for the second option for type stability?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 :)

kshyatt
kshyatt previously approved these changes Sep 24, 2026

@kshyatt kshyatt left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for catching this!

@kshyatt
kshyatt enabled auto-merge (squash) September 24, 2026 13:55
Jutho
Jutho previously approved these changes Sep 24, 2026

@Jutho Jutho left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl Outdated
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
@leburgel
leburgel dismissed stale reviews from Jutho and kshyatt via ff61cca September 24, 2026 16:27
@kshyatt

kshyatt commented Sep 25, 2026

Copy link
Copy Markdown
Member

ROCm queue is very backed up (16h) and will not be affected by this so I will just merge

@kshyatt
kshyatt disabled auto-merge September 25, 2026 05:52
@kshyatt
kshyatt merged commit 806406f into main Sep 25, 2026
11 of 12 checks passed
@kshyatt
kshyatt deleted the lb/enzyme-tensoradd-beta branch September 25, 2026 05:52
@lkdvos lkdvos mentioned this pull request Sep 30, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants