diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index 9c00e6ae8..fbb216427 100644 --- a/docs/src/Changelog.md +++ b/docs/src/Changelog.md @@ -22,8 +22,11 @@ When releasing a new version, move the "Unreleased" changes to a new version sec ### Added +- `planarcontract!` now supports an arbitrary (cyclic) output permutation `pAB`, in the same way as the non-planar `tensorcontract!`, along with the new helper `TensorKit.planar_contract_indices`, which `@planar` uses to emit canonical index tuples. ([#531](https://github.com/QuantumKitHub/TensorKit.jl/pull/531)) + ### Changed - For sector types with `GenericUnit` such that colorings are not unique, `GradedSpace`, `ProductSpace` and `HomSpace` now check for this compatibility. In particular, this prevents the construction of `TensorMap`s with incompatible colorings, which previously either errored or produced empty tensors inconsistently. ([#515](https://github.com/QuantumKitHub/TensorKit.jl/pull/515)) +- `TensorOperations.tensorcontract_structure` now requires the index tuples `pA` and `pB` to be planar (cyclic) partitions for sector types with `GenericUnit()`; hand-written planar kernels should canonicalize their index tuples with `TensorKit.planar_contract_indices` before allocating. ([#531](https://github.com/QuantumKitHub/TensorKit.jl/pull/531)) - Index manipulations use a single kernel operating on subblocks addressed by position: `StridedSubblocks` (sector-independent views into the flat data of a `TensorMap`) or `SubblockIterator` (any `AbstractTensorMap`, through `subblock`). The `TreeTransformer`s store the mapping between subblock positions and recoupling coefficients and are cached for every tensor type; conjugated and adjoint operands are handled through this mechanism instead of through `AdjointTensorMap` wrappers (internal) ([#516](https://github.com/QuantumKitHub/TensorKit.jl/issues/516)) diff --git a/docs/src/man/contractions.md b/docs/src/man/contractions.md index dcaf13bf5..8c0474c17 100644 --- a/docs/src/man/contractions.md +++ b/docs/src/man/contractions.md @@ -289,6 +289,11 @@ Finally, the name `τ` is reserved for the braiding tensor: every literal crossi The `BraidingTensor` itself does not need to be constructed by the user; the macro figures out the appropriate spaces from the surrounding contraction. Any layout the macro cannot identify as planar is rejected at parse time with `ArgumentError("not a planar diagram expression: ...")`. +For users writing their own planar kernels, it is important to note that the index tuples `pA = (oindA, cindA)` and `pB = (cindB, oindB)` handed to `planarcontract!` need not individually be planar partitions of the operands. +As long as the overall diagram is planar, the function `TensorKit.planar_contract_indices(A, pA, B, pB, pAB)` will return the canonical `pA′`, `pB′` and the remaining output permutation `pAB′` for which the operations are each planar. +These canonical tuples are also the only ones for which the intermediate spaces exist for sector types with multiple units (i.e. multifusion categories with `GenericUnit()`). +The `@planar` macro automatically applies it to the tuples it emits, but hand-written planar kernels should do the same before allocating a destination with `TensorOperations.tensoralloc_contract`. + To make this concrete, consider the contraction `A * B` for two anyonic tensors, written in a manifestly planar fashion: ```@example anyoncontraction diff --git a/src/planar/macros.jl b/src/planar/macros.jl index 1fc30ea06..a14695812 100644 --- a/src/planar/macros.jl +++ b/src/planar/macros.jl @@ -19,6 +19,9 @@ function planarparser(planarexpr, kwargs...) temporaries = Vector{Symbol}() push!(parser.postprocessors, ex -> _annotate_temporaries(ex, temporaries)) push!(parser.postprocessors, ex -> _free_temporaries(ex, temporaries)) + + partitions = Dict{Any, Union{Nothing, IndexPartition}}() + push!(parser.postprocessors, ex -> canonicalizeplanarindices(ex, partitions)) push!(parser.postprocessors, _insert_planar_operations) # braiding tensors need to be instantiated before kwargs are processed @@ -90,7 +93,8 @@ function planarparser(planarexpr, kwargs...) ) parser.contractioncostcheck = nothing push!(parser.preprocessors, ex -> _check_planarity(ex)) - push!(parser.preprocessors, ex -> _decompose_planar_contractions(ex, temporaries)) + push!(parser.preprocessors, ex -> _decompose_planar_contractions!(ex, temporaries)) + push!(parser.preprocessors, ex -> _record_index_partitions!(ex, partitions)) return parser end diff --git a/src/planar/planaroperations.jl b/src/planar/planaroperations.jl index 59a968568..5866f7ef9 100644 --- a/src/planar/planaroperations.jl +++ b/src/planar/planaroperations.jl @@ -162,15 +162,11 @@ function planarcontract!( end @timeit_debug GLOBAL_TIMER "planarcontract!" begin - codA, domA = codomainind(A), domainind(A) - codB, domB = codomainind(B), domainind(B) - oindA, cindA = pA - cindB, oindB = pB - oindA, cindA, oindB, cindB = reorder_indices( - codA, domA, codB, domB, oindA, cindA, oindB, cindB, pAB... - ) - - if oindA == codA && cindA == domA + (oindA, cindA), (cindB, oindB), pAB′ = planar_contract_indices(A, pA, B, pB, pAB) + A_in_layout = (oindA, cindA) == (codomainind(A), domainind(A)) + B_in_layout = (cindB, oindB) == (codomainind(B), domainind(B)) + + if A_in_layout A′ = A else A′ = @timeit_debug GLOBAL_TIMER "alloc: buffers" TO.tensoralloc_add( @@ -179,76 +175,113 @@ function planarcontract!( transpose!(A′, A, (oindA, cindA), One(), Zero(), backend, allocator) end - if cindB == codB && oindB == domB + if B_in_layout B′ = B else - B′ = @timeit_debug GLOBAL_TIMER "alloc: buffers" TensorOperations.tensoralloc_add( + B′ = @timeit_debug GLOBAL_TIMER "alloc: buffers" TO.tensoralloc_add( scalartype(B), B, (cindB, oindB), false, Val(true), allocator ) transpose!(B′, B, (cindB, oindB), One(), Zero(), backend, allocator) end - mul!(C, A′, B′, α, β) - (oindA == codA && cindA == domA) || TO.tensorfree!(A′, allocator) - (cindB == codB && oindB == domB) || TO.tensorfree!(B′, allocator) + + if _isdirectoutput(pAB′, length(oindA)) + mul!(C, A′, B′, α, β) + else # as in `blas_contract!`, a non-trivial `pAB` requires an intermediate + AB = @timeit_debug GLOBAL_TIMER "alloc: buffers" TO.tensoralloc_contract( + scalartype(C), A′, (codomainind(A′), domainind(A′)), false, + B′, (codomainind(B′), domainind(B′)), false, + TO.trivialpermutation(length(oindA), length(oindB)), Val(true), allocator + ) + mul!(AB, A′, B′, One(), Zero()) + transpose!(C, AB, pAB′, α, β, backend, allocator) + TO.tensorfree!(AB, allocator) + end + + A_in_layout || TO.tensorfree!(A′, allocator) + B_in_layout || TO.tensorfree!(B′, allocator) end return C end # auxiliary routines -_cyclicpermute(t::Tuple) = (Base.tail(t)..., t[1]) -_cyclicpermute(t::Tuple{}) = () - -function reorder_indices(codA, domA, codB, domB, oindA, oindB, p1, p2) - N₁ = length(oindA) - N₂ = length(oindB) - @assert length(p1) == N₁ && all(in(p1), 1:N₁) - @assert length(p2) == N₂ && all(in(p2), N₁ .+ (1:N₂)) - oindA2 = TupleTools.getindices(oindA, p1) - oindB2 = TupleTools.getindices(oindB, p2 .- N₁) - indA = (codA..., reverse(domA)...) - indB = (codB..., reverse(domB)...) - # cycle indA to be of the form (oindA2..., reverse(cindA2)...) - while length(oindA2) > 0 && indA[1] != oindA2[1] - indA = _cyclicpermute(indA) - end - # cycle indB to be of the form (cindB2..., reverse(oindB2)...) - while length(oindB2) > 0 && indB[end] != oindB2[1] - indB = _cyclicpermute(indB) - end - for i in 2:N₁ - @assert indA[i] == oindA2[i] - end - for j in 2:N₂ - @assert indB[end + 1 - j] == oindB2[j] +# whether a contraction with `N₁` open indices on `A` directly yields the destination +function _isdirectoutput(pAB::Index2Tuple, N₁::Int) + return length(pAB[1]) == N₁ && pAB == TO.trivialpermutation(pAB) +end + +# rotate `t` such that `x` comes first +_rotate_to(t::Tuple, x) = TupleTools.circshift(t, 1 - something(findfirst(==(x), t))) + +# rotate the cycle `indx` into `(head′..., reverse(tail′)...)` +function _planar_rotate(indx::IndexTuple, head::IndexTuple, tail::IndexTuple) + N₁, N = length(head), length(indx) + # rotate the arc of `head` indices in front; note that `indx` cannot be reassigned + # without boxing it in the closures below + rot = if 0 < N₁ < N + i = findfirst(ntuple(n -> indx[n] ∈ head && indx[mod1(n - 1, N)] ∉ head, Val(N))) + TupleTools.circshift(indx, 1 - something(i)) + else + indx end - Nc = length(indA) - N₁ - @assert Nc == length(indB) - N₂ - pc = ntuple(identity, Nc) - cindA2 = reverse(TupleTools.getindices(indA, N₁ .+ pc)) - cindB2 = TupleTools.getindices(indB, pc) - return oindA2, cindA2, oindB2, cindB2 + head′ = ntuple(n -> rot[n], Val(N₁)) + tail′ = reverse(ntuple(n -> rot[N₁ + n], Val(length(tail)))) + TupleTools.sort(head′) == TupleTools.sort(head) || + throw(ArgumentError(lazy"$head and $tail do not partition the cycle $indx planarly")) + return head′, tail′ end -function reorder_indices(codA, domA, codB, domB, oindA, cindA, oindB, cindB, p1, p2) - oindA2, cindA2, oindB2, cindB2 = reorder_indices( - codA, domA, codB, domB, oindA, oindB, p1, p2 +""" + IndexPartition(numout, numin) + +Stand-in for a tensor with `numout` outgoing and `numin` incoming indices, to canonicalize +index tuples with [`planar_contract_indices`](@ref) at macro-expansion time. +""" +struct IndexPartition + numout::Int + numin::Int +end +numout(p::IndexPartition) = p.numout +numin(p::IndexPartition) = p.numin + +""" + planar_contract_indices(A, pA, B, pB, pAB) -> pA′, pB′, pAB′ + +Bring the index tuples of a planar contraction into canonical form, such that `pA′` and `pB′` +are cyclic partitions of the indices of `A` and `B`, i.e. such that + + C = transpose(transpose(A, pA′) * transpose(B, pB′), pAB′) + +For sector types with `GenericUnit()` these are the only partitions with valid intermediate +spaces, so all space computations should use them. `A` and `B` can be anything supporting +`codomainind` and `domainind`, in particular `AbstractTensorMap`s, `HomSpace`s and +[`IndexPartition`](@ref)s. + +See also [`planarcontract!`](@ref). +""" +function planar_contract_indices( + A, (oindA, cindA)::Index2Tuple, + B, (cindB, oindB)::Index2Tuple, + pAB::Index2Tuple ) + indA = (codomainind(A)..., reverse(domainind(A))...) + indB = (codomainind(B)..., reverse(domainind(B))...) + oindA′, cindA′ = _planar_rotate(indA, oindA, cindA) + cindB′, oindB′ = _planar_rotate(indB, cindB, oindB) - #if oindA or oindB are empty, then reorder indices can only order it correctly up to a cyclic permutation! - if isempty(oindA2) && !isempty(cindA) - # isempty(cindA) is a cornercase which I'm not sure if we can encounter - hit = cindA[findfirst(==(first(cindB2)), cindB)] - while hit != first(cindA2) - cindA2 = _cyclicpermute(cindA2) - end - end - if isempty(oindB2) && !isempty(cindB) - hit = cindB[findfirst(==(first(cindA2)), cindA)] - while hit != first(cindB2) - cindB2 = _cyclicpermute(cindB2) - end + # if all indices are contracted, fix the residual rotation using the other tensor + if isempty(oindA′) && !isempty(cindA) + cindA′ = _rotate_to(cindA′, cindA[something(findfirst(==(first(cindB′)), cindB))]) + elseif isempty(oindB′) && !isempty(cindB) + cindB′ = _rotate_to(cindB′, cindB[something(findfirst(==(first(cindA′)), cindA))]) end - @assert TupleTools.sort(cindA) == TupleTools.sort(cindA2) - @assert TupleTools.sort(tuple.(cindA2, cindB2)) == TupleTools.sort(tuple.(cindA, cindB)) - return oindA2, cindA2, oindB2, cindB2 + TupleTools.sort(tuple.(cindA′, cindB′)) == TupleTools.sort(tuple.(cindA, cindB)) || + throw(ArgumentError(lazy"contraction of $cindA with $cindB is not planar")) + + # re-express `pAB` in terms of the reordered open indices + remap = ( + map(something, TupleTools.indexin(oindA, oindA′))..., + (length(oindA) .+ map(something, TupleTools.indexin(oindB, oindB′)))..., + ) + pAB′ = (TupleTools.getindices(remap, pAB[1]), TupleTools.getindices(remap, pAB[2])) + return (oindA′, cindA′), (cindB′, oindB′), pAB′ end diff --git a/src/planar/postprocessors.jl b/src/planar/postprocessors.jl index 8cf55ea2f..b21270c59 100644 --- a/src/planar/postprocessors.jl +++ b/src/planar/postprocessors.jl @@ -1,6 +1,6 @@ # Additional postprocessors for @planar and @plansor -# Temporaries were explicitly created by _decompose_planar_contractions and were thus +# Temporaries were explicitly created by _decompose_planar_contractions! and were thus # instantiated as if they were new output tensors rather than temporary tensors; we need # to correct for this by adding the `istemp = true` flag. function _annotate_temporaries(ex, temporaries) @@ -54,6 +54,10 @@ end # TODO: replace _planarmethod with planarmethod in everything below const _PLANAR_OPERATIONS = (:planaradd!, :planartrace!, :planarcontract!) +# `TO.instantiate` emits these with the same argument layout, i.e. with the index tuples of +# the contraction at positions 3 (`A`), 4 (`pA`), 6 (`B`), 7 (`pB`) and 9 (`pAB`) +const _CONTRACTION_INSTANTIATIONS = (:tensorcontract!, :tensoralloc_contract) + function _insert_planar_operations(ex) if isexpr(ex, :call) if ex.args[1] == GlobalRef(TensorOperations, :tensoradd!) @@ -90,6 +94,30 @@ function _insert_planar_operations(ex) return ex end +""" + canonicalizeplanarindices(ex, partitions) + +Replace the index tuples of every contraction in `ex` by their canonical form, as obtained +from [`planar_contract_indices`](@ref) and the index partitions in `partitions`. Contractions +whose operands have no known partition are left untouched. +""" +function canonicalizeplanarindices(ex, partitions) + if isexpr(ex, :call) && length(ex.args) ≥ 9 && ex.args[1] isa GlobalRef && + ex.args[1].mod === TensorOperations && + ex.args[1].name ∈ _CONTRACTION_INSTANTIATIONS + A, pA, B, pB, pAB = ex.args[3], ex.args[4], ex.args[6], ex.args[7], ex.args[9] + WA, WB = get(partitions, A, nothing), get(partitions, B, nothing) + if !isnothing(WA) && !isnothing(WB) && pA isa Index2Tuple && + pB isa Index2Tuple && pAB isa Index2Tuple + args = copy(ex.args) + args[4], args[7], args[9] = planar_contract_indices(WA, pA, WB, pB, pAB) + return Expr(:call, args...) + end + end + return ex isa Expr ? + Expr(ex.head, (canonicalizeplanarindices(a, partitions) for a in ex.args)...) : ex +end + # like `TO.insertargument`, but matching `GlobalRef`s into `TensorKit` function _insertargument(ex, arg, methods) if isexpr(ex, :call) && ex.args[1] isa GlobalRef && diff --git a/src/planar/preprocessors.jl b/src/planar/preprocessors.jl index e6b67eb64..17a99eb73 100644 --- a/src/planar/preprocessors.jl +++ b/src/planar/preprocessors.jl @@ -75,6 +75,20 @@ function _extract_tensormap_objects(ex) ) return Expr(:block, pre, pre2, ex, post) end +# used by `@planar`: record the index partition of every tensor in `ex`, keyed by its object. +# Objects that occur with conflicting partitions are recorded as `nothing`. Note that +# `_extract_tensormap_objects` checks these partitions against the actual tensors at runtime. +function _record_index_partitions!(ex, partitions) + if TO.istensor(ex) + obj, leftind, rightind = TO.decomposetensor(ex) + p = IndexPartition(length(leftind), length(rightind)) + partitions[obj] = get(partitions, obj, p) == p ? p : nothing + elseif ex isa Expr + foreach(a -> _record_index_partitions!(a, partitions), ex.args) + end + return ex +end + _is_adjoint(ex) = isexpr(ex, TO.prime) _remove_adjoint(ex) = _is_adjoint(ex) ? ex.args[1] : ex _add_adjoint(ex) = Expr(TO.prime, ex) @@ -448,8 +462,8 @@ end # decompose contraction trees in order to fix index order of temporaries # to ensure that planarity is guaranteed -_decompose_planar_contractions(ex, temporaries) = ex -function _decompose_planar_contractions(ex::Expr, temporaries) +_decompose_planar_contractions!(ex, temporaries) = ex +function _decompose_planar_contractions!(ex::Expr, temporaries) if isexpr(ex, :macrocall) && ex.args[1] == Symbol("@notensor") return ex end @@ -470,7 +484,7 @@ function _decompose_planar_contractions(ex::Expr, temporaries) end if isexpr(ex, :block) return Expr( - ex.head, [_decompose_planar_contractions(a, temporaries) for a in ex.args]... + ex.head, [_decompose_planar_contractions!(a, temporaries) for a in ex.args]... ) end return ex diff --git a/src/spaces/homspace.jl b/src/spaces/homspace.jl index 7f3418944..b41bd45be 100644 --- a/src/spaces/homspace.jl +++ b/src/spaces/homspace.jl @@ -302,25 +302,6 @@ function compose(W::HomSpace{S}, V::HomSpace{S}) where {S} return HomSpace(codomain(W), domain(V)) end -# workaround to permuting after composing intermediate spaces without constructing the latter -function _contractedspace( - A::HomSpace{S}, (oindA, cindA)::Index2Tuple, - B::HomSpace{S}, (cindB, oindB)::Index2Tuple, - (p₁, p₂)::Index2Tuple{N₁, N₂} - ) where {S, N₁, N₂} - NA = length(oindA) - - Acind = map(n -> dual(A[n]), cindA) - Bcind = map(n -> B[n], cindB) - Acind == Bcind || throw(SpaceMismatch(lazy"$(Acind) ≠ $(Bcind)")) - - getopen(n) = n <= NA ? A[oindA[n]] : B[oindB[n - NA]] - - cod = ProductSpace{S, N₁}(map(getopen, p₁)) - dom = ProductSpace{S, N₂}(map(n -> dual(getopen(n)), p₂)) - return cod ← dom -end - function TensorOperations.tensorcontract( A::HomSpace, pA::Index2Tuple, conjA::Bool, B::HomSpace, pB::Index2Tuple, conjB::Bool, @@ -341,7 +322,7 @@ function TensorOperations.tensorcontract( pB′ = adjointtensorindices(B, pB) TensorOperations.tensorcontract(A, pA, false, B′, pB′, false, pAB) else - _contractedspace(A, pA, B, pB, pAB) + permute(compose(permute(A, pA), permute(B, pB)), pAB) end end diff --git a/src/tensors/braidingtensor.jl b/src/tensors/braidingtensor.jl index cdf2f887b..0635f7988 100644 --- a/src/tensors/braidingtensor.jl +++ b/src/tensors/braidingtensor.jl @@ -211,13 +211,18 @@ function planarcontract!( length.(pA) == (2, 2) || return planarcontract!(C, TensorMap(A), pA, B, pB, pAB, α, β, backend, allocator) - spacecheck_contract(C, A, pA, false, B, pB, false, pAB) + pA′, pB′, pAB′ = planar_contract_indices(A, pA, B, pB, pAB) + # a destination that needs an additional transposition is left to the generic + # implementation: the braid below would resolve the cyclic move as crossings + _isdirectoutput(pAB′, length(pA′[1])) || + return planarcontract!(C, TensorMap(A), pA, B, pB, pAB, α, β, backend, allocator) + + spacecheck_contract(C, A, pA′, false, B, pB′, false, pAB′) codA, domA = codomainind(A), domainind(A) codB, domB = codomainind(B), domainind(B) - oindA, cindA, oindB, cindB = reorder_indices( - codA, domA, codB, domB, pA..., reverse(pB)..., pAB... - ) + oindA, cindA = pA′ + cindB, oindB = pB′ I = sectortype(C) BraidingStyle(I) isa Bosonic && @@ -264,13 +269,18 @@ function planarcontract!( length.(pB) == (2, 2) || return planarcontract!(C, A, pA, TensorMap(B), pB, pAB, α, β, backend, allocator) - spacecheck_contract(C, A, pA, false, B, pB, false, pAB) + pA′, pB′, pAB′ = planar_contract_indices(A, pA, B, pB, pAB) + # a destination that needs an additional transposition is left to the generic + # implementation: the braid below would resolve the cyclic move as crossings + _isdirectoutput(pAB′, length(pA′[1])) || + return planarcontract!(C, A, pA, TensorMap(B), pB, pAB, α, β, backend, allocator) + + spacecheck_contract(C, A, pA′, false, B, pB′, false, pAB′) codA, domA = codomainind(A), domainind(A) codB, domB = codomainind(B), domainind(B) - oindA, cindA, oindB, cindB = reorder_indices( - codA, domA, codB, domB, pA..., reverse(pB)..., pAB... - ) + oindA, cindA = pA′ + cindB, oindB = pB′ I = sectortype(C) BraidingStyle(I) isa Bosonic && diff --git a/src/tensors/linalg.jl b/src/tensors/linalg.jl index b93f20cb3..657c4649a 100644 --- a/src/tensors/linalg.jl +++ b/src/tensors/linalg.jl @@ -326,12 +326,19 @@ function LinearAlgebra.tr(t::AbstractTensorMap) return s end +# LinearAlgebra's BLAS wrappers do not accept `VectorInterface.One`/`Zero`, e.g. +# `herk_wrapper!` (hit by `mul!(C, A, A')`) calls `isreal` on the scalars. +_blasscalar(α::Number) = α +_blasscalar(::One) = true +_blasscalar(::Zero) = false + # TensorMap multiplication function LinearAlgebra.mul!( tC::AbstractTensorMap, tA::AbstractTensorMap, tB::AbstractTensorMap, α = true, β = false ) compose(space(tA), space(tB)) == space(tC) || throw(SpaceMismatch(lazy"$(space(tC)) ≠ $(space(tA)) * $(space(tB))")) + α, β = _blasscalar(α), _blasscalar(β) @timeit_debug GLOBAL_TIMER "dense: matmul" begin iterC = blocks(tC) @@ -602,7 +609,7 @@ is `domain(t1) ⊗ domain(t2)`. function ⊗(A::AbstractTensorMap, B::AbstractTensorMap) check_spacetype(A, B) - # allocate destination with correct scalartype + # index tuples for the blockwise tensor product pA = ((codomainind(A)..., domainind(A)...), ()) pB = ((), (codomainind(B)..., domainind(B)...)) NA = numind(A) @@ -610,8 +617,12 @@ function ⊗(A::AbstractTensorMap, B::AbstractTensorMap) (codomainind(A)..., (codomainind(B) .+ NA)...), (domainind(A)..., (domainind(B) .+ NA)...), ) + # note that we don't use `tensoralloc_contract`: its intermediate spaces are not + # cyclically ordered, which is not allowed for `GenericUnit` sectors TC = TO.promote_contract(scalartype(A), scalartype(B)) - C = TO.tensoralloc_contract(TC, A, pA, false, B, pB, false, pAB, Val(false)) + TTC = TO.tensorcontract_type(TC, A, pA, false, B, pB, false, pAB) + structure = (codomain(A) ⊗ codomain(B)) ← (domain(A) ⊗ domain(B)) + C = TO.tensoralloc(TTC, structure, Val(false)) zerovector!(C) # implement tensor product diff --git a/src/tensors/tensoroperations.jl b/src/tensors/tensoroperations.jl index 91fa672bb..14e803c47 100644 --- a/src/tensors/tensoroperations.jl +++ b/src/tensors/tensoroperations.jl @@ -167,9 +167,9 @@ function TO.tensorcontract_structure( B::AbstractTensorMap, pB::Index2Tuple, conjB::Bool, pAB::Index2Tuple{N₁, N₂} ) where {N₁, N₂} - VA, pA′ = conjA ? (space(A)', adjointtensorindices(A, pA)) : (space(A), pA) - VB, pB′ = conjB ? (space(B)', adjointtensorindices(B, pB)) : (space(B), pB) - return _contractedspace(VA, pA′, VB, pB′, pAB) + sA = TO.tensoradd_structure(A, pA, conjA) + sB = TO.tensoradd_structure(B, pB, conjB) + return permute(compose(sA, sB), pAB) end function TO.checkcontractible( diff --git a/test/tensors/contractions.jl b/test/tensors/contractions.jl index e457c39d5..1132809dc 100644 --- a/test/tensors/contractions.jl +++ b/test/tensors/contractions.jl @@ -57,6 +57,28 @@ for V in spacelist @planar t5[a; b] := t4[a c; b c] @test t2 ≈ t5 end + @timedtestset "Planar contraction: test self-consistency" begin + t = rand(ComplexF64, V1 ⊗ V2 ⊗ V3 ← (V4 ⊗ V5)') + # neither index partition of this contraction is planar by itself, + # only their cyclic rotations are + @planar ρ[a; b] := t[a c d; e f] * t'[e f; b c d] + @test space(ρ) == (V1 ← V1) + @test ρ ≈ ρ' + t1 = transpose(t, ((1,), (4, 5, 3, 2))) + t2 = transpose(t', ((1, 2, 5, 4), (3,))) + @test ρ ≈ t1 * t2 + if BraidingStyle(I) isa Bosonic + # `@planar` and `@tensor` only agree for bosonic braiding + @tensor ρ2[a; b] := t[a c d; e f] * conj(t[b c d; e f]) + @test ρ ≈ ρ2 + end + # the intermediate result is allocated as a temporary + @planar ρ3[a; b] := t[a c d; e f] * t'[e f; g c d] * ρ[g; b] + @test ρ3 ≈ ρ * ρ + # the same, but with a less trivial temporary + @planar ρ4[a; b] := t[a c d; e f] * (t'[e f; g c d] * ρ[g; b]) + @test ρ4 ≈ ρ * ρ + end if BraidingStyle(I) isa Bosonic && hasfusiontensor(I) @timedtestset "Trace: test via conversion" begin t = rand(ComplexF64, V1 ⊗ V2' ⊗ V3 ⊗ V2 ⊗ V1' ⊗ V3') @@ -97,6 +119,12 @@ for V in spacelist t2 = rand(T, V2 ⊗ V3, V4') t = @constinferred (t1 ⊗ t2) @test norm(t) ≈ norm(t1) * norm(t2) + # a factor with more than one index in the domain has a non-planar + # intermediate space + t2′ = transpose(t2, ((1,), (3, 2))) + t′ = @constinferred (t1 ⊗ t2′) + @test norm(t′) ≈ norm(t1) * norm(t2′) + @test space(t′) == (V1 ⊗ V2 ← V5' ⊗ V4' ⊗ V3') end end if BraidingStyle(I) isa Bosonic && hasfusiontensor(I) diff --git a/test/tensors/planar.jl b/test/tensors/planar.jl index b5d0bbd45..72a933b21 100644 --- a/test/tensors/planar.jl +++ b/test/tensors/planar.jl @@ -4,6 +4,7 @@ using TensorKit using TensorKit: type_repr using TensorKit: PlanarTrivial, ℙ using TensorKit: planaradd!, planartrace!, planarcontract! +using TensorKit: planar_contract_indices, SpaceMismatch using TensorOperations spacelist = default_spacelist(fast_tests) @@ -95,6 +96,64 @@ end @test force_planar(tensorcontract!(C, A, pA, false, B, pB, false, pAB, true, true)) ≈ planarcontract!(C′, A′, pA, B′, pB, pAB, true, true) + + # an output permutation that is not absorbed by the cyclic reordering + pAB2 = ((2, 1), (3, 4, 5)) + D = randn((ℂ^2)' ⊗ ℂ^2 ← ℂ^5 ⊗ (ℂ^2)' ⊗ ℂ^4) + D′ = force_planar(D) + @test force_planar(tensorcontract!(D, A, pA, false, B, pB, false, pAB2, true, true)) ≈ + planarcontract!(D′, A′, pA, B′, pB, pAB2, true, true) + + # an output permutation that is not cyclic is not planar + pAB3 = ((1, 2), (3, 4, 5)) + E′ = force_planar(randn(ℂ^2 ⊗ (ℂ^2)' ← ℂ^5 ⊗ (ℂ^2)' ⊗ ℂ^4)) + @test_throws ArgumentError planarcontract!(E′, A′, pA, B′, pB, pAB3, true, true) + end + + @testset "planar_contract_indices" begin + V1, V2, V3, V4, V5 = VIBM + W = V1 ⊗ V2 ⊗ V3 ← (V4 ⊗ V5)' + pA, pB = ((1,), (2, 3, 4, 5)), ((4, 5, 1, 2), (3,)) + pAB = ((1,), (2,)) + + # the partitions of a planar contraction need not be planar by themselves + @test_throws SpaceMismatch permute(W, pA) + @test_throws SpaceMismatch permute(W', pB) + + pA′, pB′, pAB′ = @constinferred planar_contract_indices(W, pA, W', pB, pAB) + @test permute(W, pA′) isa TensorKit.HomSpace + @test permute(W', pB′) isa TensorKit.HomSpace + @test TensorOperations.tensorcontract(W, pA′, false, W', pB′, false, pAB′) == + (V1 ← V1) + + # all indices contracted: the rotations are only fixed by the other factor + pA0, pB0 = ((), (1, 2, 3, 4, 5)), ((3, 4, 5, 1, 2), ()) + pA0′, pB0′, pAB0′ = @constinferred planar_contract_indices( + W, pA0, W', pB0, ((), ()) + ) + @test permute(W, pA0′) isa TensorKit.HomSpace + @test permute(W', pB0′) isa TensorKit.HomSpace + @test numind( + TensorOperations.tensorcontract(W, pA0′, false, W', pB0′, false, pAB0′) + ) == 0 + + # not a planar contraction + @test_throws ArgumentError planar_contract_indices( + W, ((1,), (3, 2, 4, 5)), W', pB, pAB + ) + @test_throws ArgumentError planar_contract_indices( + W, ((2,), (1, 3, 4, 5)), W', pB, pAB + ) + + # the output permutation is remapped along with the reordered open indices + WA = ℂ^2 ⊗ ℂ^3 ← ℂ^2 ⊗ ℂ^5 ⊗ ℂ^4 + WB = ℂ^2 ⊗ ℂ^4 ← ℂ^4 ⊗ ℂ^3 + pA2, pB2 = ((1, 3, 4), (5, 2)), ((2, 4), (1, 3)) + pA2′, pB2′, pAB2′ = planar_contract_indices(WA, pA2, WB, pB2, ((3, 2, 1), (4, 5))) + @test (pA2′, pB2′) == (((4, 3, 1), (5, 2)), ((2, 4), (1, 3))) + @test pAB2′ == ((1, 2, 3), (4, 5)) + @test last(planar_contract_indices(WA, pA2, WB, pB2, ((2, 1), (3, 4, 5)))) == + ((2, 3), (1, 4, 5)) end end @@ -102,32 +161,49 @@ end T = ComplexF64 @testset "backend and allocator insertion" begin - # trailing arguments of every `planar*!` call in `ex` - function planartrailing(ex, out = Any[]) + PLANAR_OPERATIONS = (:planaradd!, :planartrace!, :planarcontract!) + # trailing arguments of every call in `ex` whose name is in `names` + function planartrailing(ex, names, out = Any[]) ex isa Expr || return out if Meta.isexpr(ex, :call) && ex.args[1] isa GlobalRef && - ex.args[1].name in (:planaradd!, :planartrace!, :planarcontract!) + ex.args[1].name in names push!(out, ex.args[end]) end - foreach(a -> planartrailing(a, out), ex.args) + foreach(a -> planartrailing(a, names, out), ex.args) return out end ex = @macroexpand @planar backend = MarkerBackend() C[i; j] := A[i; k l] * τ[k l; m n] * B[m n; j] - trailing = planartrailing(ex) + trailing = planartrailing(ex, PLANAR_OPERATIONS) @test !isempty(trailing) @test all(==(:(MarkerBackend())), trailing) # an allocator implies a default backend, and both land on the planar calls ex = @macroexpand @planar allocator = MarkerAllocator() C[i; j] := A[i; k l] * τ[k l; m n] * B[m n; j] - trailing = planartrailing(ex) + trailing = planartrailing(ex, PLANAR_OPERATIONS) @test !isempty(trailing) @test all(==(:(MarkerAllocator())), trailing) @test occursin("DefaultBackend", string(ex)) end + @testset "canonical index tuples" begin + # the emitted partitions are planar, unlike the raw ones of the decomposition + function planarindices(ex, out = Any[]) + ex isa Expr || return out + if Meta.isexpr(ex, :call) && ex.args[1] isa GlobalRef && + ex.args[1].name === :planarcontract! + push!(out, (ex.args[4], ex.args[6], ex.args[7])) + end + foreach(a -> planarindices(a, out), ex.args) + return out + end + ex = @macroexpand @planar ρ[a; b] := t[a c d; e f] * u[e f; b c d] + @test planarindices(ex) == + [(((1,), (4, 5, 3, 2)), ((1, 2, 5, 4), (3,)), ((1,), (2,)))] + end + @testset "allocator is rewound" begin # A `BufferAllocator` hands out slices of a single buffer and reclaims them only # by rewinding its offset -- `tensorfree!` is a no-op for it. The temporaries a