Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 17 additions & 13 deletions ext/TensorOperationsChainRulesCoreExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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...
Expand All @@ -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)
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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...
Expand Down
5 changes: 5 additions & 0 deletions test/ad.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ using TensorOperations
using TensorOperations: StridedBLAS, StridedNative
using Test
using ChainRulesTestUtils
using ChainRulesCore: NoTangent
using VectorInterface: Zero

ChainRulesTestUtils.test_method_tables()

Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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();
Expand Down
Loading