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
31 changes: 18 additions & 13 deletions src/precompile/contract.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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
Expand Down
12 changes: 10 additions & 2 deletions src/precompile/indexmanipulations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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₁))
Expand All @@ -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
Expand Down
3 changes: 1 addition & 2 deletions src/precompile/precompile.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading