diff --git a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl index 82bd9b578..af6f35125 100644 --- a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl +++ b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl @@ -12,6 +12,8 @@ function ChainRulesCore.rrule( α::Number, β::Number, 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 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 @@ -66,6 +68,7 @@ function ChainRulesCore.rrule( α::Number, β::Number, ba... ) C′ = tensorcontract!(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(), @@ -156,6 +159,7 @@ function ChainRulesCore.rrule( α::Number, β::Number, ba... ) C′ = tensortrace!(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