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
3 changes: 3 additions & 0 deletions docs/src/Changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -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))

Expand Down
5 changes: 5 additions & 0 deletions docs/src/man/contractions.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 5 additions & 1 deletion src/planar/macros.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
161 changes: 97 additions & 64 deletions src/planar/planaroperations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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)))
Comment thread
Jutho marked this conversation as resolved.

# 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
30 changes: 29 additions & 1 deletion src/planar/postprocessors.jl
Original file line number Diff line number Diff line change
@@ -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)
Expand Down Expand Up @@ -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!)
Expand Down Expand Up @@ -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 &&
Expand Down
20 changes: 17 additions & 3 deletions src/planar/preprocessors.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
21 changes: 1 addition & 20 deletions src/spaces/homspace.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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

Expand Down
Loading
Loading