From 2c1df65d455377290579922346bf505a2fd7b128 Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 24 Sep 2026 14:14:07 +0200 Subject: [PATCH 1/3] =?UTF-8?q?Fix=20Enzyme=20`tensoradd!`=20rule=20cachin?= =?UTF-8?q?g=20`C`=20based=20on=20the=20value=20of=20`=CE=B2`=20instead=20?= =?UTF-8?q?of=20its=20activity?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../TensorOperationsEnzymeExt.jl | 2 +- test/enzyme.jl | 9 +++++++++ 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl index 4afb5292..77fe514e 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) : nothing 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 From 869e7d3e5e8af9bb68e8644566961026e378375d Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 24 Sep 2026 15:43:25 +0200 Subject: [PATCH 2/3] Preserve type of second option --- 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 77fe514e..e429a412 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 = !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 From ff61ccaaddb6219870ec0c21dabb483399d73fb3 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Thu, 24 Sep 2026 18:27:24 +0200 Subject: [PATCH 3/3] Update ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl Co-authored-by: Jutho --- 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 e429a412..4f314d06 100644 --- a/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl +++ b/ext/TensorOperationsEnzymeExt/TensorOperationsEnzymeExt.jl @@ -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