diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index 4afb5292..4f314d06 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -189,7 +189,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 = !iszero(β_dβ.val) ? copy(C_dC.val) : C_dC.val + cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : C_dC.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) α = α_dα.val β = β_dβ.val @@ -309,7 +309,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) : nothing + cache_C = !isa(β_dβ, Const) ? copy(C_dC.val) : C_dC.val ba = map(ba_ -> getfield(ba_, :val), ba_dba) α = α_dα.val β = β_dβ.val diff --git a/test/enzyme.jl b/test/enzyme.jl index 16612c8d..aa7245cf 100644 --- a/test/enzyme.jl +++ b/test/enzyme.jl @@ -251,3 +251,12 @@ end test_reverse(tensorscalar, Active, (C, Duplicated); atol, rtol) test_forward(tensorscalar, Duplicated, (C, Duplicated); atol, rtol) end + +# dβ needs the original C whenever β is Active, also when its value is zero +@testset "tensoradd! with Active β = 0" begin + pA = ((2, 1, 4, 3, 5), ()) + A = rand(Float64, (2, 3, 4, 2, 1)) + C = rand(Float64, size.(Ref(A), pA[1])) + 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