Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 38 additions & 23 deletions ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)},
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -97,33 +112,33 @@ 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)
zero(α_dα.val)
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)
zero(β_dβ.val)
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

Expand All @@ -148,18 +163,18 @@ 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)
end
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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -222,28 +237,28 @@ 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)
zero(α_dα.val)
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)
zero(β_dβ.val)
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

Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -344,28 +359,28 @@ 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)
zero(α_dα.val)
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)
zero(β_dβ.val)
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

Expand All @@ -391,15 +406,15 @@ 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)
end
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
Expand Down
25 changes: 25 additions & 0 deletions test/enzyme.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading