diff --git a/src/precompile/contract.jl b/src/precompile/contract.jl index 74097bd35..44c08220d 100644 --- a/src/precompile/contract.jl +++ b/src/precompile/contract.jl @@ -20,6 +20,7 @@ function precompile_contract(::Type{S}; eltypes = PRECOMPILE_ELTYPES, ndims = PR V = unitspace(S) backend = TO.DefaultBackend() allocator = TO.DefaultAllocator() + symmetric_braiding = BraidingStyle(sectortype(S)) isa SymmetricBraiding for T in eltypes α, β = rand(T), rand(T) @@ -32,27 +33,31 @@ function precompile_contract(::Type{S}; eltypes = PRECOMPILE_ELTYPES, ndims = PR B = randn(T, W ← V) pA = ((1,), ntuple(i -> i + 1, Val(N - 1))) pB = (ntuple(identity, Val(N - 1)), (N,)) - C = TO.tensoralloc_contract(T, A, pA, false, B, pB, false, ((1,), (2,)), Val(false)) - # both scalar-type paths: generic `(α,β)` and the identity `(One(), Zero())` fast path - TO.tensorcontract!(C, A, pA, false, B, pB, false, ((1,), (2,)), α, β, backend, allocator) - TO.tensorcontract!(C, A, pA, false, B, pB, false, ((1,), (2,)), One(), Zero(), backend, allocator) - - # a non-trivial permutation exercises the repartition/braid/transform machinery at - # arity N (the contraction above takes the no-copy view path for its natural partition) - permute(A, (ntuple(i -> N - i + 1, Val(N)), ())) + pAB = ((1,), (2,)) + C = TO.tensoralloc_contract(T, A, pA, false, B, pB, false, pAB, Val(false)) + + planarcontract!(C, A, pA, B, pB, pAB, α, β, backend, allocator) + planarcontract!(C, A, pA, B, pB, pAB, One(), Zero(), backend, allocator) + + if symmetric_braiding + TO.tensorcontract!(C, A, pA, false, B, pB, false, pAB, α, β, backend, allocator) + TO.tensorcontract!(C, A, pA, false, B, pB, false, pAB, One(), Zero(), backend, allocator) + end end # the conjugated-operand branch (`conjA=true`) is a distinct runtime path (adjoint - # handling) that `conj(A) * B` networks hit; `@tensor` handles the space bookkeeping + # handling) that `conj(A) * B` networks hit; `@plansor` picks the planar or non-planar + # implementation depending on the braiding style, so this compiles for any sectortype A2 = randn(T, V ← V) B2 = randn(T, V ← V) - @tensor Cc[a; c] := conj(A2[b; a]) * B2[b; c] + @plansor Cc[a; c] := conj(A2[b; a]) * B2[b; c] # partial trace (the two traced legs are mutually dual) At = randn(T, V ⊗ V' ← V) - TO.tensortrace!( - TO.tensoralloc_add(T, At, ((3,), ()), false, Val(false)), - At, ((3,), ()), ((1,), (2,)), false, α, β, backend, allocator + Ct = TO.tensoralloc_add(T, At, ((3,), ()), false, Val(false)) + planartrace!(Ct, At, ((3,), ()), ((1,), (2,)), α, β, backend, allocator) + symmetric_braiding && TO.tensortrace!( + Ct, At, ((3,), ()), ((1,), (2,)), false, α, β, backend, allocator ) end return nothing diff --git a/src/precompile/indexmanipulations.jl b/src/precompile/indexmanipulations.jl index 4e97d1b20..efa3fed97 100644 --- a/src/precompile/indexmanipulations.jl +++ b/src/precompile/indexmanipulations.jl @@ -17,6 +17,8 @@ See also [`precompile_contract`](@ref TensorKit.precompile_contract), [`precompi """ function precompile_indexmanipulations(::Type{S}; eltypes = PRECOMPILE_ELTYPES, ndims = PRECOMPILE_NDIMS) where {S <: IndexSpace} V = unitspace(S) + symmetric_braiding = BraidingStyle(sectortype(S)) isa SymmetricBraiding + has_braiding = !(BraidingStyle(sectortype(S)) isa NoBraiding) for T in eltypes # `Val(N)`/`ntuple` keep the index tuples concrete so the machinery specializes per arity for N in 1:ndims @@ -30,8 +32,13 @@ function precompile_indexmanipulations(::Type{S}; eltypes = PRECOMPILE_ELTYPES, p2 = ntuple(i -> perm[i + N₁], Val(N - N₁)) # `permute` and `braid` funnel through `add_transform!` with different transformers - permute(t, (p1, p2)) - braid(t, (p1, p2), ntuple(identity, Val(N))) # `levels` is a tuple over the source indices + if symmetric_braiding + permute(t, (p1, p2)) + permute(t', (p1, p2)) + elseif has_braiding + braid(t, (p1, p2), ntuple(identity, Val(N))) # `levels` is a tuple over the source indices + braid(t', (p1, p2), ntuple(identity, Val(N))) # `levels` is a tuple over the source indices + end # canonical transpose `(reverse(domain), reverse(codomain))` is always a valid cyclic transpose tp1 = ntuple(i -> N - i + 1, Val(N - N₁)) @@ -42,6 +49,7 @@ function precompile_indexmanipulations(::Type{S}; eltypes = PRECOMPILE_ELTYPES, repartition(t, N₁ < N ? N₁ + 1 : N₁ - 1) twist(t, 1) + twist(t', 1) end end return nothing diff --git a/src/precompile/precompile.jl b/src/precompile/precompile.jl index bf03b54c4..918bb1b05 100644 --- a/src/precompile/precompile.jl +++ b/src/precompile/precompile.jl @@ -3,9 +3,8 @@ module Precompilation export precompile_indexmanipulations, precompile_contract, precompile_factorizations using ..TensorKit -using ..TensorKit: TO +using ..TensorKit: TO, planarcontract!, planartrace! using VectorInterface: One, Zero -using TensorOperations: @tensor using PrecompileTools: @setup_workload, @compile_workload using Preferences: @load_preference