diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index 4f314d06..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) && !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) && !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) && !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) && !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 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