From 905aeafa76c72874e99498244380009564dad147 Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 24 Sep 2026 12:26:22 +0200 Subject: [PATCH 1/2] =?UTF-8?q?Avoid=20copying=20and=20retaining=20`C`=20i?= =?UTF-8?q?n=20TensorOperations=20rrules=20when=20`=CE=B2=20=3D=20Zero()`?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../tensoroperations.jl | 16 ++++++++++------ test/chainrules/tensoroperations.jl | 7 +++++++ 2 files changed, 17 insertions(+), 6 deletions(-) diff --git a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl index 82bd9b578..6ff1c0ddb 100644 --- a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl +++ b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl @@ -11,7 +11,9 @@ function ChainRulesCore.rrule( A::AbstractTensorMap, pA::Index2Tuple, conjA::Bool, α::Number, β::Number, ba... ) - C′ = tensoradd!(copy(C), A, pA, conjA, α, β, ba...) + C′ = tensoradd!(β === Zero() ? similar(C) : copy(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) @@ -49,7 +51,7 @@ function ChainRulesCore.rrule( else ZeroTangent() end - dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C, ΔC))) : ZeroTangent() + dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C_β, ΔC))) : NoTangent() dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, dA, NoTangent(), NoTangent(), dα, dβ, dba... end @@ -65,7 +67,8 @@ function ChainRulesCore.rrule( pAB::Index2Tuple, α::Number, β::Number, ba... ) - C′ = tensorcontract!(copy(C), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) + C′ = tensorcontract!(β === Zero() ? similar(C) : copy(C), A, pA, conjA, B, pB, conjB, pAB, α, β, ba...) + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectB = ProjectTo(B) @@ -138,7 +141,7 @@ function ChainRulesCore.rrule( else ZeroTangent() end - dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C, ΔC))) : ZeroTangent() + dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C_β, ΔC))) : NoTangent() dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, dA, NoTangent(), NoTangent(), @@ -155,7 +158,8 @@ function ChainRulesCore.rrule( A::AbstractTensorMap, p::Index2Tuple, q::Index2Tuple, conjA::Bool, α::Number, β::Number, ba... ) - C′ = tensortrace!(copy(C), A, p, q, conjA, α, β, ba...) + C′ = tensortrace!(β === Zero() ? similar(C) : copy(C), A, p, q, conjA, α, β, ba...) + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectC = ProjectTo(C) @@ -190,7 +194,7 @@ function ChainRulesCore.rrule( else ZeroTangent() end - dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C, ΔC))) : ZeroTangent() + dβ = _needs_tangent(β) ? @thunk(projectβ(inner(C_β, ΔC))) : NoTangent() dba = map(_ -> NoTangent(), ba) return NoTangent(), dC, dA, NoTangent(), NoTangent(), NoTangent(), dα, dβ, dba... end diff --git a/test/chainrules/tensoroperations.jl b/test/chainrules/tensoroperations.jl index 13052bca6..c3fd557b8 100644 --- a/test/chainrules/tensoroperations.jl +++ b/test/chainrules/tensoroperations.jl @@ -3,6 +3,7 @@ using TensorKit using TensorKit: type_repr, SectorDict using TensorOperations using ChainRulesCore +using VectorInterface: Zero using ChainRulesTestUtils using Random using LinearAlgebra @@ -47,6 +48,7 @@ for V in spacelist for conjA in (false, true) C = randn!(TensorOperations.tensoralloc_add(T, A, p, conjA, Val(false))) test_rrule(tensortrace!, C, A, p, q, conjA, α, β; atol, rtol) + test_rrule(tensortrace!, C, A, p, q, conjA, α, Zero() ⊢ NoTangent(); atol, rtol) end end end @@ -65,6 +67,7 @@ for V in spacelist C2 = randn!(TensorOperations.tensoralloc_add(T, A, p, true, Val(false))) test_rrule(tensoradd!, C2, A, p, true, α, β; atol, rtol) + test_rrule(tensoradd!, C1, A, p, false, α, Zero() ⊢ NoTangent(); atol, rtol) A = rand(Bool) ? C1 : C2 end @@ -107,6 +110,10 @@ for V in spacelist tensorcontract!, C, A, pA, conjA, B, pB, conjB, pAB, α, β; atol, rtol ) + test_rrule( + tensorcontract!, C, A, pA, conjA, B, pB, conjB, pAB, α, + Zero() ⊢ NoTangent(); atol, rtol + ) end end end From d82411f70a948c9f1cfa8f157e5215e97c193dfa Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 25 Sep 2026 08:33:09 +0200 Subject: [PATCH 2/2] Simplify --- ext/TensorKitChainRulesCoreExt/tensoroperations.jl | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl index 6ff1c0ddb..af6f35125 100644 --- a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl +++ b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl @@ -11,7 +11,7 @@ function ChainRulesCore.rrule( A::AbstractTensorMap, pA::Index2Tuple, conjA::Bool, α::Number, β::Number, ba... ) - C′ = tensoradd!(β === Zero() ? similar(C) : copy(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 @@ -67,7 +67,7 @@ function ChainRulesCore.rrule( pAB::Index2Tuple, α::Number, β::Number, ba... ) - C′ = tensorcontract!(β === Zero() ? similar(C) : copy(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) @@ -158,7 +158,7 @@ function ChainRulesCore.rrule( A::AbstractTensorMap, p::Index2Tuple, q::Index2Tuple, conjA::Bool, α::Number, β::Number, ba... ) - C′ = tensortrace!(β === Zero() ? similar(C) : copy(C), A, p, q, conjA, α, β, ba...) + C′ = tensortrace!(copy(C), A, p, q, conjA, α, β, ba...) C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A)