From 35b10f6c1fe948ec4105a6e8a4c70328b4ed7bf2 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 23 Sep 2026 13:48:42 +0200 Subject: [PATCH 1/5] Protect from writing into dval if it's === val --- .../TensorOperationsEnzymeExt.jl | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index 4f314d06..43b46805 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -97,12 +97,12 @@ function EnzymeRules.reverse( β = β_dβ.val pA, pB, pAB, conjA, conjB = getfield.((pA_dpA, pB_dpB, pAB_dpAB, conjA_dconjA, conjB_dconjB), :val) - if !isa(A_dA, Const) && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensorcontract_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) end - if !isa(B_dB, Const) && !isa(C_dC, Const) + if !isa(B_dB, Const) && B_dB.dval !== B.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔB = B_dB.dval TensorOperations.tensorcontract_pullback_dB!(ΔB, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) @@ -222,7 +222,7 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensoradd_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, α, ba...) @@ -344,7 +344,7 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensortrace_pullback_dA!(ΔA, ΔC, Cval, Aval, p, q, conjA, α, ba...) From 4ef966db5a6a365bd7b0684b7274b4e0a60be28a Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 23 Sep 2026 13:56:56 +0200 Subject: [PATCH 2/5] Typo --- ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index 43b46805..e0af9065 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -97,7 +97,7 @@ function EnzymeRules.reverse( β = β_dβ.val pA, pB, pAB, conjA, conjB = getfield.((pA_dpA, pB_dpB, pAB_dpAB, conjA_dconjA, conjB_dconjB), :val) - if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensorcontract_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) @@ -222,7 +222,7 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensoradd_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, α, ba...) @@ -344,7 +344,7 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && A_dA.dval !== A.val && !isa(C_dC, Const) + if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensortrace_pullback_dA!(ΔA, ΔC, Cval, Aval, p, q, conjA, α, ba...) From d6af75f6ff5245c0e3b4b44976828352b53f8edf Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 23 Sep 2026 13:57:19 +0200 Subject: [PATCH 3/5] Typo again --- ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index e0af9065..f76a69c9 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -102,7 +102,7 @@ function EnzymeRules.reverse( ΔA = A_dA.dval TensorOperations.tensorcontract_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) end - if !isa(B_dB, Const) && B_dB.dval !== B.val && !isa(C_dC, Const) + if !isa(B_dB, Const) && B_dB.dval !== B_dB.val && !isa(C_dC, Const) ΔC = C_dC.dval ΔB = B_dB.dval TensorOperations.tensorcontract_pullback_dB!(ΔB, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) From 5920d12c4ed6feec478f439ac557cd0cd381ca1d Mon Sep 17 00:00:00 2001 From: lkdvos Date: Wed, 30 Sep 2026 10:25:00 -0400 Subject: [PATCH 4/5] add helper function and guard cache copy(C) if not needed due to runtime activity --- .../TensorOperationsEnzymeExt.jl | 61 ++++++++++++------- 1 file changed, 38 insertions(+), 23 deletions(-) diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index f76a69c9..8094b908 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -16,6 +16,21 @@ using Enzyme.EnzymeCore: EnzymeRules @inline EnzymeRules.inactive_type(v::Type{<:Index2Tuple}) = true @inline EnzymeRules.inactive_type(v::Type{<:IndexTuple}) = true +""" + is_inactive(config, x::Annotation) -> Bool + +Check whether the annotated argument `x` should be treated as inactive, i.e. whether no tangent should be read from or accumulated into `x.dval`. +Under runtime activity, Enzyme may pass an argument that is inactive at run time as a `Duplicated` whose shadow aliases the primal (`x.dval === x.val`); +writing into that shadow would then corrupt the primal. +See https://enzymead.github.io/Enzyme.jl/dev/faq/#faq-runtime-activity. + +!!! note + This is a stopgap until an equivalent helper is available in `EnzymeRules` + (see https://github.com/EnzymeAD/Enzyme.jl/pull/3597#discussion_r4014313289). +""" +@inline is_inactive(config, ::Const) = true +@inline is_inactive(config, x::Annotation) = EnzymeRules.runtime_activity(config) && x.dval === x.val + function EnzymeRules.augmented_primal( config::EnzymeRules.RevConfigWidth{1}, func::Const{typeof(TensorOperations.tensoralloc)}, @@ -62,7 +77,7 @@ function EnzymeRules.augmented_primal( # form caches if needed cache_A = EnzymeRules.overwritten(config)[3] ? copy(A_dA.val) : nothing cache_B = EnzymeRules.overwritten(config)[6] ? copy(B_dB.val) : nothing - cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : C_dC.val + cache_C = !isa(β_dβ, Const) && !is_inactive(config, C_dC) ? copy(C_dC.val) : C_dC.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) TensorOperations.tensorcontract!(C_dC.val, A_dA.val, pA_dpA.val, conjA_dconjA.val, B_dB.val, pB_dpB.val, conjB_dconjB.val, pAB_dpAB.val, α_dα.val, β_dβ.val, ba...) primal = EnzymeRules.needs_primal(config) ? C_dC.val : nothing @@ -97,17 +112,17 @@ function EnzymeRules.reverse( β = β_dβ.val pA, pB, pAB, conjA, conjB = getfield.((pA_dpA, pB_dpB, pAB_dpAB, conjA_dconjA, conjB_dconjB), :val) - if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) + if !is_inactive(config, A_dA) && !is_inactive(config, C_dC) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensorcontract_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) end - if !isa(B_dB, Const) && B_dB.dval !== B_dB.val && !isa(C_dC, Const) + if !is_inactive(config, B_dB) && !is_inactive(config, C_dC) ΔC = C_dC.dval ΔB = B_dB.dval TensorOperations.tensorcontract_pullback_dB!(ΔB, ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) end - Δα = if !isa(α_dα, Const) && !isa(C_dC, Const) + Δα = if !isa(α_dα, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.tensorcontract_pullback_dα(ΔC, Cval, Aval, pA, conjA, Bval, pB, conjB, pAB, α, ba...) elseif !isa(α_dα, Const) @@ -115,7 +130,7 @@ function EnzymeRules.reverse( else nothing end - Δβ = if !isa(β_dβ, Const) && !isa(C_dC, Const) + Δβ = if !isa(β_dβ, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.pullback_dβ(ΔC, Cval, β) elseif !isa(β_dβ, Const) @@ -123,7 +138,7 @@ function EnzymeRules.reverse( else nothing end - !isa(C_dC, Const) && TensorOperations.pullback_dC!(C_dC.dval, β) + !is_inactive(config, C_dC) && TensorOperations.pullback_dC!(C_dC.dval, β) return nothing, nothing, nothing, nothing, nothing, nothing, nothing, nothing, Δα, Δβ, map(ba_ -> nothing, ba)... end @@ -148,7 +163,7 @@ function EnzymeRules.forward( β = β_dβ.val pA, pB, pAB, conjA, conjB = getfield.((pA_dpA, pB_dpB, pAB_dpAB, conjA_dconjA, conjB_dconjB), :val) - if !isa(C_dC, Const) + if !is_inactive(config, C_dC) scale!(C_dC.dval, β) if !isa(β_dβ, Const) add!(C_dC.dval, C_dC.val, β_dβ.dval) @@ -156,10 +171,10 @@ function EnzymeRules.forward( if !isa(α_dα, Const) tensorcontract!(C_dC.dval, A_dA.val, pA, conjA, B_dB.val, pB, conjB, pAB, α_dα.dval, One(), ba...) end - if !isa(A_dA, Const) + if !is_inactive(config, A_dA) tensorcontract!(C_dC.dval, A_dA.dval, pA, conjA, B_dB.val, pB, conjB, pAB, α, One(), ba...) end - if !isa(B_dB, Const) + if !is_inactive(config, B_dB) tensorcontract!(C_dC.dval, A_dA.val, pA, conjA, B_dB.dval, pB, conjB, pAB, α, One(), ba...) end end @@ -189,7 +204,7 @@ function EnzymeRules.augmented_primal( ) where {RT, Tα <: Number, Tβ <: Number, TA <: Number, TC <: Number} # form caches if needed cache_A = EnzymeRules.overwritten(config)[3] ? copy(A_dA.val) : nothing - cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : C_dC.val + cache_C = !isa(β_dβ, Const) && !is_inactive(config, C_dC) ? copy(C_dC.val) : C_dC.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) α = α_dα.val β = β_dβ.val @@ -222,12 +237,12 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) + if !is_inactive(config, A_dA) && !is_inactive(config, C_dC) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensoradd_pullback_dA!(ΔA, ΔC, Cval, Aval, pA, conjA, α, ba...) end - Δα = if !isa(α_dα, Const) && !isa(C_dC, Const) + Δα = if !isa(α_dα, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.tensoradd_pullback_dα(ΔC, Cval, Aval, pA, conjA, α, ba...) elseif !isa(α_dα, Const) @@ -235,7 +250,7 @@ function EnzymeRules.reverse( else nothing end - Δβ = if !isa(β_dβ, Const) && !isa(C_dC, Const) + Δβ = if !isa(β_dβ, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.pullback_dβ(ΔC, Cval, β) elseif !isa(β_dβ, Const) @@ -243,7 +258,7 @@ function EnzymeRules.reverse( else nothing end - !isa(C_dC, Const) && TensorOperations.pullback_dC!(C_dC.dval, β) + !is_inactive(config, C_dC) && TensorOperations.pullback_dC!(C_dC.dval, β) return nothing, nothing, nothing, nothing, Δα, Δβ, map(ba_ -> nothing, ba)... end @@ -270,12 +285,12 @@ function EnzymeRules.forward( # dD = dα * A + α * dA + β dC + dβ * C # dC′ = β dC + dβ * C - if !isa(C_dC, Const) + if !is_inactive(config, C_dC) scale!(C_dC.dval, β) if !isa(β_dβ, Const) add!(C_dC.dval, C_dC.val, β_dβ.dval) end - if !isa(A_dA, Const) + if !is_inactive(config, A_dA) TensorOperations.tensoradd!(C_dC.dval, A_dA.dval, pA, conjA, α, One(), ba...) end if !isa(α_dα, Const) @@ -309,7 +324,7 @@ function EnzymeRules.augmented_primal( ) where {RT, Tα <: Number, Tβ <: Number, TA <: Number, TC <: Number} # form caches if needed cache_A = EnzymeRules.overwritten(config)[3] ? copy(A_dA.val) : nothing - cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : C_dC.val + cache_C = !isa(β_dβ, Const) && !is_inactive(config, C_dC) ? copy(C_dC.val) : C_dC.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) α = α_dα.val β = β_dβ.val @@ -344,12 +359,12 @@ function EnzymeRules.reverse( β = β_dβ.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) - if !isa(A_dA, Const) && A_dA.dval !== A_dA.val && !isa(C_dC, Const) + if !is_inactive(config, A_dA) && !is_inactive(config, C_dC) ΔC = C_dC.dval ΔA = A_dA.dval TensorOperations.tensortrace_pullback_dA!(ΔA, ΔC, Cval, Aval, p, q, conjA, α, ba...) end - Δα = if !isa(α_dα, Const) && !isa(C_dC, Const) + Δα = if !isa(α_dα, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.tensortrace_pullback_dα(ΔC, Cval, Aval, p, q, conjA, α, ba...) elseif !isa(α_dα, Const) @@ -357,7 +372,7 @@ function EnzymeRules.reverse( else nothing end - Δβ = if !isa(β_dβ, Const) && !isa(C_dC, Const) + Δβ = if !isa(β_dβ, Const) && !is_inactive(config, C_dC) ΔC = C_dC.dval TensorOperations.pullback_dβ(ΔC, Cval, β) elseif !isa(β_dβ, Const) @@ -365,7 +380,7 @@ function EnzymeRules.reverse( else nothing end - !isa(C_dC, Const) && TensorOperations.pullback_dC!(C_dC.dval, β) + !is_inactive(config, C_dC) && TensorOperations.pullback_dC!(C_dC.dval, β) return nothing, nothing, nothing, nothing, nothing, Δα, Δβ, map(ba_ -> nothing, ba)... end @@ -391,7 +406,7 @@ function EnzymeRules.forward( # dD = dα * tr(A) + α * tr(dA) + dβ * C + β * dC # dC1 = dβ * C + β * dC - if !isa(C_dC, Const) + if !is_inactive(config, C_dC) scale!(C_dC.dval, β) if !isa(β_dβ, Const) add!(C_dC.dval, C_dC.val, β_dβ.dval) @@ -399,7 +414,7 @@ function EnzymeRules.forward( if !isa(α_dα, Const) TensorOperations.tensortrace!(C_dC.dval, A_dA.val, p, q, conjA, α_dα.dval, One(), ba...) end - if !isa(A_dA, Const) + if !is_inactive(config, A_dA) TensorOperations.tensortrace!(C_dC.dval, A_dA.dval, p, q, conjA, α, One(), ba...) end end From 8596efff2fb395dd65956080722e346d4f86d6ac Mon Sep 17 00:00:00 2001 From: lkdvos Date: Wed, 30 Sep 2026 10:26:15 -0400 Subject: [PATCH 5/5] Add `is_inactive` helper for runtime activity in Enzyme rules MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Under runtime activity, Enzyme can pass an argument that is inactive at run time as a `Duplicated` whose shadow aliases the primal. Replace the inline `dval !== val` guards with a single helper, `is_inactive(config, x)`, which checks `x isa Const || (runtime_activity(config) && x.dval === x.val)`, and apply it to all tensor arguments (A, B, C) in the forward and reverse rules of `tensorcontract!`, `tensoradd!` and `tensortrace!`. This also skips copying `C` for `Δβ` when `C` is inactive. The helper is meant to be replaced by an upstream EnzymeRules equivalent once available. Add a regression test adapted from EnzymeAD/Enzyme.jl#3623. Co-Authored-By: Claude Opus 5.5 --- test/enzyme.jl | 25 +++++++++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/test/enzyme.jl b/test/enzyme.jl index aa7245cf..1f6064c6 100644 --- a/test/enzyme.jl +++ b/test/enzyme.jl @@ -260,3 +260,28 @@ end test_reverse(tensoradd!, Duplicated, (C, Duplicated), (A, Duplicated), (pA, Const), (false, Const), (randn(), Active), (0.0, Active)) test_reverse(tensoradd!, Duplicated, (C, Duplicated), (A, Duplicated), (pA, Const), (false, Const), (randn(), Active), (randn(), Const)) end + +# under runtime activity, inactive operands can arrive as `Duplicated` with `dval === val`, +# see https://github.com/EnzymeAD/Enzyme.jl/issues/3623 +const RUNTIME_ACTIVITY_OPS = Dict([1, 2] => randn(4, 4), [2, 3] => randn(4, 4)) +function runtime_activity_f(x) + total = 0.0 + for (_, op) in collect(RUNTIME_ACTIVITY_OPS) + C = zeros(length(x)) + @tensor C[i] += op[i, j] * x[j] + D = x * x' + @tensor D[j, i] += op[i, j] + T = fill(x[1]) + @tensor T[] += op[i, i] + total += sum(C) + sum(D) + T[] + end + return total +end +@testset "runtime activity with inactive operands" begin + ops = deepcopy(RUNTIME_ACTIVITY_OPS) + x = randn(4) + dx = zero(x) + Enzyme.autodiff(set_runtime_activity(Reverse), Const(runtime_activity_f), Active, Duplicated(x, dx)) + @test all(RUNTIME_ACTIVITY_OPS[k] == ops[k] for k in keys(ops)) + @test dx ≈ sum(vec(sum(op; dims = 1)) .+ 2 * sum(x) .+ [1, 0, 0, 0] for op in values(ops)) +end