diff --git a/ext/TensorOperationsChainRulesCoreExt.jl b/ext/TensorOperationsChainRulesCoreExt.jl index 4ea1123f..1ca4c721 100644 --- a/ext/TensorOperationsChainRulesCoreExt.jl +++ b/ext/TensorOperationsChainRulesCoreExt.jl @@ -71,6 +71,8 @@ function ChainRulesCore.rrule( end function _rrule_tensoradd!(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 projectA = ProjectTo(A) projectC = ProjectTo(C) @@ -81,16 +83,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... @@ -112,6 +114,7 @@ function ChainRulesCore.rrule( 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_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectB = ProjectTo(B) @@ -123,17 +126,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, @@ -156,6 +159,7 @@ function ChainRulesCore.rrule( end function _rrule_tensortrace!(C, A, p, q, conjA, α, β, ba) C′ = tensortrace!(copy(C), A, p, q, conjA, α, β, ba...) + C_β = _needs_tangent(β) ? C : nothing projectA = ProjectTo(A) projectC = ProjectTo(C) @@ -166,16 +170,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..1be00fc8 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: NoTangent +using VectorInterface: Zero 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();