diff --git a/test/enzyme-indexmanipulations-flip-twist-transform/flip.jl b/test/enzyme-indexmanipulations-flip-twist-transform/flip.jl index e2de670fc..66efa44f4 100644 --- a/test/enzyme-indexmanipulations-flip-twist-transform/flip.jl +++ b/test/enzyme-indexmanipulations-flip-twist-transform/flip.jl @@ -6,22 +6,42 @@ using Random spacelist = ad_spacelist(fast_tests) eltypes = (Float64, ComplexF64) -@timedtestset "Enzyme - Index Manipulations (flip):" begin +is_ci = get(ENV, "CI", "false") == "true" + +@timedtestset "Enzyme - Index Manipulations (flip and twist):" begin @timedtestset "$(TensorKit.type_repr(sectortype(eltype(V)))) ($T) TA ($TA)" for V in spacelist, T in eltypes, TA in (Duplicated,) atol = default_tol(T) rtol = default_tol(T) has_braiding = BraidingStyle(sectortype(eltype(V))) isa HasBraiding if has_braiding A = randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])') + if !(T <: Real && !(sectorscalartype(sectortype(A)) <: Real)) + EnzymeTestUtils.test_reverse(twist!, TA, (A, TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) + EnzymeTestUtils.test_reverse(twist!, TA, (A, TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) + if !is_ci + EnzymeTestUtils.test_reverse(twist!, TA, (A, TA), (1, Const); atol, rtol) + EnzymeTestUtils.test_reverse(twist!, TA, (A, TA), ([1, 3], Const); atol, rtol) + end + EnzymeTestUtils.test_forward(twist!, TA, (A, TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) + EnzymeTestUtils.test_forward(twist!, TA, (A, TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) + if !is_ci + EnzymeTestUtils.test_forward(twist!, TA, (A, TA), (1, Const); atol, rtol) + EnzymeTestUtils.test_forward(twist!, TA, (A, TA), ([1, 3], Const); atol, rtol) + end + end EnzymeTestUtils.test_reverse(flip, TA, (A, TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) EnzymeTestUtils.test_reverse(flip, TA, (A, TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) - EnzymeTestUtils.test_reverse(flip, TA, (A, TA), (1, Const); atol, rtol) - EnzymeTestUtils.test_reverse(flip, TA, (A, TA), ([1, 3], Const); atol, rtol) + if !is_ci + EnzymeTestUtils.test_reverse(flip, TA, (A, TA), (1, Const); atol, rtol) + EnzymeTestUtils.test_reverse(flip, TA, (A, TA), ([1, 3], Const); atol, rtol) + end EnzymeTestUtils.test_forward(flip, TA, (A, TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) EnzymeTestUtils.test_forward(flip, TA, (A, TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) - EnzymeTestUtils.test_forward(flip, TA, (A, TA), (1, Const); atol, rtol) - EnzymeTestUtils.test_forward(flip, TA, (A, TA), ([1, 3], Const); atol, rtol) + if !is_ci + EnzymeTestUtils.test_forward(flip, TA, (A, TA), (1, Const); atol, rtol) + EnzymeTestUtils.test_forward(flip, TA, (A, TA), ([1, 3], Const); atol, rtol) + end end end end diff --git a/test/enzyme-indexmanipulations-flip-twist-transform/permute.jl b/test/enzyme-indexmanipulations-flip-twist-transform/permute.jl index 28102a3b1..e61036677 100644 --- a/test/enzyme-indexmanipulations-flip-twist-transform/permute.jl +++ b/test/enzyme-indexmanipulations-flip-twist-transform/permute.jl @@ -12,7 +12,6 @@ Tαs = is_ci ? (Active,) : (Active, Const) Tβs = is_ci ? (Active,) : (Active, Const) @timedtestset "Enzyme - Index Manipulations (permute!): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes - println(TensorKit.type_repr(sectortype(eltype(V)))) atol = default_tol(T) rtol = default_tol(T) symmetricbraiding = BraidingStyle(sectortype(eltype(V))) isa SymmetricBraiding diff --git a/test/enzyme-indexmanipulations-flip-twist-transform/twist.jl b/test/enzyme-indexmanipulations-flip-twist-transform/twist.jl deleted file mode 100644 index 2cacc2b3d..000000000 --- a/test/enzyme-indexmanipulations-flip-twist-transform/twist.jl +++ /dev/null @@ -1,28 +0,0 @@ -using Test, TestExtras -using TensorKit -using TensorOperations -using VectorInterface: Zero, One -using Enzyme, EnzymeTestUtils -using Random - -spacelist = ad_spacelist(fast_tests) -eltypes = (Float64, ComplexF64) - -@timedtestset "Enzyme - Index Manipulations (twist):" begin - @timedtestset "$(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, TA in (Duplicated,) - atol = default_tol(T) - rtol = default_tol(T) - A = randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])') - has_braiding = BraidingStyle(sectortype(eltype(V))) isa HasBraiding - if has_braiding && !(T <: Real && !(sectorscalartype(sectortype(A)) <: Real)) - EnzymeTestUtils.test_reverse(twist!, TA, (copy(A), TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) - EnzymeTestUtils.test_reverse(twist!, TA, (copy(A), TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) - EnzymeTestUtils.test_reverse(twist!, TA, (copy(A), TA), (1, Const); atol, rtol) - EnzymeTestUtils.test_reverse(twist!, TA, (copy(A), TA), ([1, 3], Const); atol, rtol) - EnzymeTestUtils.test_forward(twist!, TA, (copy(A), TA), (1, Const); atol, rtol, fkwargs = (inv = false,)) - EnzymeTestUtils.test_forward(twist!, TA, (copy(A), TA), ([1, 3], Const); atol, rtol, fkwargs = (inv = true,)) - EnzymeTestUtils.test_forward(twist!, TA, (copy(A), TA), (1, Const); atol, rtol) - EnzymeTestUtils.test_forward(twist!, TA, (copy(A), TA), ([1, 3], Const); atol, rtol) - end - end -end