From fd412866bb7aea07db939f5ac9702558b6eb2f33 Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 24 Sep 2026 14:54:28 +0200 Subject: [PATCH 1/2] =?UTF-8?q?Avoid=20copying=20and=20retaining=20`C`=20i?= =?UTF-8?q?n=20ChainRules=20rules=20when=20`=CE=B2=20=3D=20Zero()`?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ext/TensorOperationsChainRulesCoreExt.jl | 40 ++++++++++++++---------- test/ad.jl | 12 +++++++ 2 files changed, 36 insertions(+), 16 deletions(-) diff --git a/ext/TensorOperationsChainRulesCoreExt.jl b/ext/TensorOperationsChainRulesCoreExt.jl index 4ea1123f..03bec2e4 100644 --- a/ext/TensorOperationsChainRulesCoreExt.jl +++ b/ext/TensorOperationsChainRulesCoreExt.jl @@ -58,6 +58,10 @@ function ChainRulesCore.rrule(::typeof(tensorscalar), C) return tensorscalar(C), tensorscalar_pullback end +# with β = Zero() the contents of `C` are not used, but some kernels need defined entries +_output_buffer(C, β) = copy(C) +_output_buffer(C, ::Zero) = isbitstype(scalartype(C)) ? similar(C) : zerovector!!(similar(C)) + # The current `rrule` design makes sure that the implementation for custom types does # not need to support the backend or allocator arguments function ChainRulesCore.rrule( @@ -70,7 +74,9 @@ function ChainRulesCore.rrule( return _rrule_tensoradd!(C, A, pA, conjA, α, β, ba) end function _rrule_tensoradd!(C, A, pA, conjA, α, β, ba) - C′ = tensoradd!(copy(C), A, pA, conjA, α, β, ba...) + C′ = tensoradd!(_output_buffer(C, β), A, pA, conjA, α, β, ba...) + # only keep `C` alive on the tape if `dβ` needs it + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectC = ProjectTo(C) @@ -81,16 +87,16 @@ function _rrule_tensoradd!(C, A, pA, conjA, α, β, ba) ΔC = unthunk(ΔC′) dC = β === Zero() ? ZeroTangent() : @thunk projectC(pullback_dC(ΔC, β)) - dA = @thunk projectA(tensoradd_pullback_dA(ΔC, C, A, pA, conjA, α, ba...)) + dA = @thunk projectA(tensoradd_pullback_dA(ΔC, C_β, A, pA, conjA, α, ba...)) dα = if _needs_tangent(α) - @thunk projectα(tensoradd_pullback_dα(ΔC, C, A, pA, conjA, α, ba...)) + @thunk projectα(tensoradd_pullback_dα(ΔC, C_β, A, pA, conjA, α, ba...)) else ZeroTangent() end dβ = if _needs_tangent(β) - @thunk projectβ(pullback_dβ(ΔC, C, β)) + @thunk projectβ(pullback_dβ(ΔC, C_β, β)) else - ZeroTangent() + NoTangent() end dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, dA, NoTangent(), NoTangent(), dα, dβ, dba... @@ -111,7 +117,8 @@ function ChainRulesCore.rrule( return _rrule_tensorcontract!(C, A, pA, conjA, B, pB, conjB, pAB, α, β, ba) end function _rrule_tensorcontract!(C, A, pA, conjA, B, pB, conjB, pAB, α, β, ba) - C′ = tensorcontract!(copy(C), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) + C′ = tensorcontract!(_output_buffer(C, β), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectB = ProjectTo(B) @@ -123,17 +130,17 @@ function _rrule_tensorcontract!(C, A, pA, conjA, B, pB, conjB, pAB, α, β, ba) ΔC = unthunk(ΔC′) dC = β === Zero() ? ZeroTangent() : @thunk projectC(pullback_dC(ΔC, β)) - dA = @thunk projectA(tensorcontract_pullback_dA(ΔC, C, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) - dB = @thunk projectB(tensorcontract_pullback_dB(ΔC, C, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) + dA = @thunk projectA(tensorcontract_pullback_dA(ΔC, C_β, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) + dB = @thunk projectB(tensorcontract_pullback_dB(ΔC, C_β, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) dα = if _needs_tangent(α) - @thunk projectα(tensorcontract_pullback_dα(ΔC, C, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) + @thunk projectα(tensorcontract_pullback_dα(ΔC, C_β, A, pA, conjA, B, pB, conjB, pAB, α, ba...)) else ZeroTangent() end dβ = if _needs_tangent(β) - @thunk projectβ(pullback_dβ(ΔC, C, β)) + @thunk projectβ(pullback_dβ(ΔC, C_β, β)) else - ZeroTangent() + NoTangent() end dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, @@ -155,7 +162,8 @@ function ChainRulesCore.rrule( return _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) end function _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) - C′ = tensortrace!(copy(C), A, p, q, conjA, α, β, ba...) + C′ = tensortrace!(_output_buffer(C, β), A, p, q, conjA, α, β, ba...) + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectC = ProjectTo(C) @@ -166,16 +174,16 @@ function _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) ΔC = unthunk(ΔC′) dC = β === Zero() ? ZeroTangent() : @thunk projectC(pullback_dC(ΔC, β)) - dA = @thunk projectA(tensortrace_pullback_dA(ΔC, C, A, p, q, conjA, α, ba...)) + dA = @thunk projectA(tensortrace_pullback_dA(ΔC, C_β, A, p, q, conjA, α, ba...)) dα = if _needs_tangent(α) - @thunk projectα(tensortrace_pullback_dα(ΔC, C, A, p, q, conjA, α, ba...)) + @thunk projectα(tensortrace_pullback_dα(ΔC, C_β, A, p, q, conjA, α, ba...)) else ZeroTangent() end dβ = if _needs_tangent(β) - @thunk projectβ(pullback_dβ(ΔC, C, β)) + @thunk projectβ(pullback_dβ(ΔC, C_β, β)) else - ZeroTangent() + NoTangent() end dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, dA, NoTangent(), NoTangent(), NoTangent(), dα, dβ, dba... diff --git a/test/ad.jl b/test/ad.jl index 3d8e50e1..ca862700 100644 --- a/test/ad.jl +++ b/test/ad.jl @@ -2,6 +2,8 @@ using TensorOperations using TensorOperations: StridedBLAS, StridedNative using Test using ChainRulesTestUtils +using ChainRulesCore: rrule, NoTangent +using VectorInterface: Zero, One ChainRulesTestUtils.test_method_tables() @@ -23,6 +25,7 @@ ChainRulesTestUtils.test_method_tables() test_rrule(tensortrace!, C, A, p, q, false, α, β; atol, rtol) test_rrule(tensortrace!, C, A, p, q, true, α, β; atol, rtol) + test_rrule(tensortrace!, C, A, p, q, false, α, Zero() ⊢ NoTangent(); atol, rtol) test_rrule(tensortrace!, C, A, p, q, true, α, β, StridedBLAS(); atol, rtol) test_rrule(tensortrace!, C, A, p, q, false, α, β, StridedNative(); atol, rtol) @@ -52,6 +55,7 @@ end β = rand(T) test_rrule(tensoradd!, C, A, pA, false, α, β; atol, rtol) test_rrule(tensoradd!, C, A, pA, true, α, β; atol, rtol) + test_rrule(tensoradd!, C, A, pA, true, α, Zero() ⊢ NoTangent(); atol, rtol) test_rrule(tensoradd!, C, A, pA, false, α, β, StridedBLAS(); atol, rtol) test_rrule(tensoradd!, C, A, pA, true, α, β, StridedNative(); atol, rtol) @@ -80,6 +84,7 @@ end test_rrule(tensorcontract!, C, A, pA, true, B, pB, false, pAB, α, β; atol, rtol) test_rrule(tensorcontract!, C, A, pA, false, B, pB, true, pAB, α, β; atol, rtol) test_rrule(tensorcontract!, C, A, pA, true, B, pB, true, pAB, α, β; atol, rtol) + test_rrule(tensorcontract!, C, A, pA, false, B, pB, true, pAB, α, Zero() ⊢ NoTangent(); atol, rtol) test_rrule( tensorcontract!, C, A, pA, false, B, pB, false, pAB, α, β, StridedBLAS(); @@ -99,3 +104,10 @@ end fill!(C, rand(T)) test_rrule(tensorscalar, C; atol, rtol) end + +# with β = Zero(), `C` may be uninitialized, also for non-isbits element types +@testset "β = Zero() with uninitialized C" begin + A = rand(BigFloat, 3, 4); B = rand(BigFloat, 4, 5) + C′, = rrule(tensorcontract!, similar(A, 3, 5), A, ((1,), (2,)), false, B, ((1,), (2,)), false, ((1, 2), ()), One(), Zero()) + @test C′ ≈ A * B +end From 6b933ded1b497b3c61b79769850ce95057010f86 Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 25 Sep 2026 08:20:08 +0200 Subject: [PATCH 2/2] Simplify --- ext/TensorOperationsChainRulesCoreExt.jl | 10 +++------- test/ad.jl | 11 ++--------- 2 files changed, 5 insertions(+), 16 deletions(-) diff --git a/ext/TensorOperationsChainRulesCoreExt.jl b/ext/TensorOperationsChainRulesCoreExt.jl index 03bec2e4..1ca4c721 100644 --- a/ext/TensorOperationsChainRulesCoreExt.jl +++ b/ext/TensorOperationsChainRulesCoreExt.jl @@ -58,10 +58,6 @@ function ChainRulesCore.rrule(::typeof(tensorscalar), C) return tensorscalar(C), tensorscalar_pullback end -# with β = Zero() the contents of `C` are not used, but some kernels need defined entries -_output_buffer(C, β) = copy(C) -_output_buffer(C, ::Zero) = isbitstype(scalartype(C)) ? similar(C) : zerovector!!(similar(C)) - # The current `rrule` design makes sure that the implementation for custom types does # not need to support the backend or allocator arguments function ChainRulesCore.rrule( @@ -74,7 +70,7 @@ function ChainRulesCore.rrule( return _rrule_tensoradd!(C, A, pA, conjA, α, β, ba) end function _rrule_tensoradd!(C, A, pA, conjA, α, β, ba) - C′ = tensoradd!(_output_buffer(C, β), A, pA, conjA, α, β, ba...) + C′ = tensoradd!(copy(C), A, pA, conjA, α, β, ba...) # only keep `C` alive on the tape if `dβ` needs it C_β = _needs_tangent(β) ? C : nothing @@ -117,7 +113,7 @@ function ChainRulesCore.rrule( return _rrule_tensorcontract!(C, A, pA, conjA, B, pB, conjB, pAB, α, β, ba) end function _rrule_tensorcontract!(C, A, pA, conjA, B, pB, conjB, pAB, α, β, ba) - C′ = tensorcontract!(_output_buffer(C, β), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) + C′ = tensorcontract!(copy(C), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) @@ -162,7 +158,7 @@ function ChainRulesCore.rrule( return _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) end function _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) - C′ = tensortrace!(_output_buffer(C, β), A, p, q, conjA, α, β, ba...) + C′ = tensortrace!(copy(C), A, p, q, conjA, α, β, ba...) C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) diff --git a/test/ad.jl b/test/ad.jl index ca862700..1be00fc8 100644 --- a/test/ad.jl +++ b/test/ad.jl @@ -2,8 +2,8 @@ using TensorOperations using TensorOperations: StridedBLAS, StridedNative using Test using ChainRulesTestUtils -using ChainRulesCore: rrule, NoTangent -using VectorInterface: Zero, One +using ChainRulesCore: NoTangent +using VectorInterface: Zero ChainRulesTestUtils.test_method_tables() @@ -104,10 +104,3 @@ end fill!(C, rand(T)) test_rrule(tensorscalar, C; atol, rtol) end - -# with β = Zero(), `C` may be uninitialized, also for non-isbits element types -@testset "β = Zero() with uninitialized C" begin - A = rand(BigFloat, 3, 4); B = rand(BigFloat, 4, 5) - C′, = rrule(tensorcontract!, similar(A, 3, 5), A, ((1,), (2,)), false, B, ((1,), (2,)), false, ((1, 2), ()), One(), Zero()) - @test C′ ≈ A * B -end