From 40451d64137aece6fc8cefc78b99dbc7dfa685a1 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Fri, 11 Sep 2026 14:07:58 -0400 Subject: [PATCH 1/6] Generalize `planarcontract!` to arbitrary `pAB`, remove `_contractedspace` MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `@planar` allocates through `TO.tensoralloc_contract` with the raw index tuples of the planar decomposition, in which `pA` and `pB` need not be planar partitions by themselves; only `planarcontract!` rotated them into shape, and only when `pAB` could be absorbed into those rotations. Replace `reorder_indices` by `planar_contract_indices`, which canonicalizes `pA` and `pB` and remaps `pAB` instead of absorbing it, let `planarcontract!` apply a residual cyclic `pAB` with `transpose!` through an intermediate, and allocate planar destinations through the new `planaralloc_contract`. Allocate the destination of `⊗` directly, since its intermediates are non-planar by construction. With every caller passing planar partitions, `tensorcontract_structure` can again compose the intermediate spaces, which also restores contractions of operands with different `spacetype` parameters. Co-Authored-By: Claude Opus 5 (1M context) --- docs/src/Changelog.md | 3 + docs/src/man/contractions.md | 5 + .../tensoroperations.jl | 2 + ext/TensorKitMooncakeExt/utility.jl | 1 + src/planar/planaroperations.jl | 163 +++++++++++------- src/planar/postprocessors.jl | 14 +- src/spaces/homspace.jl | 21 +-- src/tensors/braidingtensor.jl | 26 ++- src/tensors/linalg.jl | 8 +- src/tensors/tensoroperations.jl | 6 +- test/tensors/contractions.jl | 25 +++ test/tensors/planar.jl | 76 +++++++- 12 files changed, 246 insertions(+), 104 deletions(-) diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index 9c00e6ae8..dbc19829b 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 internal helpers `TensorKit.planar_contract_indices` and `TensorKit.planaralloc_contract`. ([#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()`; planar code should allocate through `TensorKit.planaralloc_contract`, which canonicalizes them. ([#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/ext/TensorKitChainRulesCoreExt/tensoroperations.jl b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl index 82bd9b578..5e695c35b 100644 --- a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl +++ b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl @@ -198,6 +198,8 @@ function ChainRulesCore.rrule( return C′, pullback end +@non_differentiable TensorKit.planaralloc_contract(args...) + function ChainRulesCore.rrule(::typeof(TensorKit.scalar), t::AbstractTensorMap) val = scalar(t) function scalar_pullback(Δval) diff --git a/ext/TensorKitMooncakeExt/utility.jl b/ext/TensorKitMooncakeExt/utility.jl index 1bedf566c..145c83e0b 100644 --- a/ext/TensorKitMooncakeExt/utility.jl +++ b/ext/TensorKitMooncakeExt/utility.jl @@ -19,6 +19,7 @@ Mooncake.tangent_type(::Type{<:HomSpace}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorstructure), AbstractTensorMap, Int, Bool} @zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorcontract_structure), AbstractTensorMap, Index2Tuple, Bool, AbstractTensorMap, Index2Tuple, Bool, Index2Tuple} +@zero_derivative DefaultCtx Tuple{typeof(TensorKit.planaralloc_contract), Any} @zero_derivative DefaultCtx Tuple{typeof(TensorKit.planar_trace), TensorKit.FusionTreePair, Index2Tuple, Index2Tuple} @zero_derivative DefaultCtx Tuple{typeof(TensorKit.has_shared_permute), AbstractTensorMap, Index2Tuple} diff --git a/src/planar/planaroperations.jl b/src/planar/planaroperations.jl index 59a968568..95b3b468d 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,117 @@ 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 +""" + 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 and `HomSpace`s. + +See also [`planarcontract!`](@ref) and [`planaralloc_contract`](@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 + # 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))]) end - if isempty(oindB2) && !isempty(cindB) - hit = cindB[findfirst(==(first(cindA2)), cindA)] - while hit != first(cindB2) - cindB2 = _cyclicpermute(cindB2) - end + if 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 + +""" + planaralloc_contract(TC, A, pA, B, pB, pAB, [istemp, allocator]) + +Allocate the destination of `planarcontract!(C, A, pA, B, pB, pAB, α, β)`. + +The planar counterpart of `TensorOperations.tensoralloc_contract`: the index tuples are +canonicalized with [`planar_contract_indices`](@ref) first, such that the space computation +only involves valid intermediate spaces. +""" +function planaralloc_contract( + TC, A, pA::Index2Tuple, B, pB::Index2Tuple, pAB::Index2Tuple, + istemp::Val = Val(false), allocator = TO.DefaultAllocator() + ) + pA′, pB′, pAB′ = planar_contract_indices(A, pA, B, pB, pAB) + return TO.tensoralloc_contract(TC, A, pA′, false, B, pB′, false, pAB′, istemp, allocator) end diff --git a/src/planar/postprocessors.jl b/src/planar/postprocessors.jl index 8cf55ea2f..72bccfcf4 100644 --- a/src/planar/postprocessors.jl +++ b/src/planar/postprocessors.jl @@ -53,6 +53,7 @@ end # TODO: replace _planarmethod with planarmethod in everything below const _PLANAR_OPERATIONS = (:planaradd!, :planartrace!, :planarcontract!) +const _PLANAR_ALLOCATIONS = (:planaralloc_contract,) function _insert_planar_operations(ex) if isexpr(ex, :call) @@ -78,6 +79,14 @@ function _insert_planar_operations(ex) ex.head, GlobalRef(TensorKit, Symbol(:planartrace!)), map(_insert_planar_operations, ex.args[2:end])... ) + elseif ex.args[1] == GlobalRef(TensorOperations, :tensoralloc_contract) + conjB = popat!(ex.args, 8) + conjA = popat!(ex.args, 5) + @assert !conjA && !conjB "conj flags should be disabled ($conjA), ($conjB)" + return Expr( + ex.head, GlobalRef(TensorKit, Symbol(:planaralloc_contract)), + map(_insert_planar_operations, ex.args[2:end])... + ) elseif ex.args[1] in TensorOperations.tensoroperationsfunctions return Expr( ex.head, GlobalRef(TensorOperations, ex.args[1]), @@ -116,10 +125,11 @@ end """ insertplanarallocator(ex, allocator) -Insert the allocator argument into the tensor operation methods `planaradd!`, `planartrace!`, and `planarcontract!`. +Insert the allocator argument into the tensor operation methods `planaradd!`, `planartrace!`, +`planarcontract!`, and `planaralloc_contract`. See also: [`TensorOperations.insertallocator`](@ref). """ function insertplanarallocator(ex, allocator) - return _insertargument(ex, allocator, _PLANAR_OPERATIONS) + return _insertargument(ex, allocator, (_PLANAR_OPERATIONS..., _PLANAR_ALLOCATIONS...)) end 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..f0b000b0e 100644 --- a/src/tensors/linalg.jl +++ b/src/tensors/linalg.jl @@ -602,7 +602,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 +610,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..195e7851e 100644 --- a/test/tensors/contractions.jl +++ b/test/tensors/contractions.jl @@ -57,6 +57,25 @@ 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 ≈ ρ * ρ + 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 +116,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..1f476e867 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,30 +161,35 @@ end T = ComplexF64 @testset "backend and allocator insertion" begin - # trailing arguments of every `planar*!` call in `ex` - function planartrailing(ex, out = Any[]) + # 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, (:planaradd!, :planartrace!, :planarcontract!)) @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, (:planaradd!, :planartrace!, :planarcontract!, :planaralloc_contract) + ) @test !isempty(trailing) @test all(==(:(MarkerAllocator())), trailing) @test occursin("DefaultBackend", string(ex)) + + alloc_trailing = planartrailing(ex, (:planaralloc_contract,)) + @test !isempty(alloc_trailing) end @testset "allocator is rewound" begin From eb7daa9f8ed1ce80ffee9c1159eb5e3868ca51f2 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Fri, 11 Sep 2026 15:38:56 -0400 Subject: [PATCH 2/6] Emit canonical planar index tuples at macro-expansion time `planar_contract_indices` only needs the index partitions of both operands, which `@planar` knows: `_extract_tensormap_objects` checks the partitions written in the expression against the actual tensors. Record them while preprocessing and canonicalize the index tuples of the emitted planar contractions, so that the runtime canonicalization is left with nothing to do for macro-generated code. Co-Authored-By: Claude Opus 5 (1M context) --- src/planar/macros.jl | 4 ++++ src/planar/planaroperations.jl | 16 +++++++++++++++- src/planar/postprocessors.jl | 24 ++++++++++++++++++++++++ src/planar/preprocessors.jl | 14 ++++++++++++++ test/tensors/planar.jl | 16 ++++++++++++++++ 5 files changed, 73 insertions(+), 1 deletion(-) diff --git a/src/planar/macros.jl b/src/planar/macros.jl index 1fc30ea06..aceb8dad8 100644 --- a/src/planar/macros.jl +++ b/src/planar/macros.jl @@ -21,6 +21,9 @@ function planarparser(planarexpr, kwargs...) push!(parser.postprocessors, ex -> _free_temporaries(ex, temporaries)) push!(parser.postprocessors, _insert_planar_operations) + partitions = Dict{Any, Union{Nothing, IndexPartition}}() + push!(parser.postprocessors, ex -> canonicalizeplanarindices(ex, partitions)) + # braiding tensors need to be instantiated before kwargs are processed push!(parser.preprocessors, _construct_braidingtensors) @@ -91,6 +94,7 @@ 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 -> _index_partitions!(partitions, ex)) return parser end diff --git a/src/planar/planaroperations.jl b/src/planar/planaroperations.jl index 95b3b468d..0ad133dcf 100644 --- a/src/planar/planaroperations.jl +++ b/src/planar/planaroperations.jl @@ -230,6 +230,19 @@ function _planar_rotate(indx::IndexTuple, head::IndexTuple, tail::IndexTuple) return head′, tail′ end +""" + 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′ @@ -240,7 +253,8 @@ are cyclic partitions of the indices of `A` and `B`, i.e. such that 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 and `HomSpace`s. +`codomainind` and `domainind`, in particular `AbstractTensorMap`s, `HomSpace`s and +[`IndexPartition`](@ref)s. See also [`planarcontract!`](@ref) and [`planaralloc_contract`](@ref). """ diff --git a/src/planar/postprocessors.jl b/src/planar/postprocessors.jl index 72bccfcf4..de5a08c08 100644 --- a/src/planar/postprocessors.jl +++ b/src/planar/postprocessors.jl @@ -54,6 +54,7 @@ end # TODO: replace _planarmethod with planarmethod in everything below const _PLANAR_OPERATIONS = (:planaradd!, :planartrace!, :planarcontract!) const _PLANAR_ALLOCATIONS = (:planaralloc_contract,) +const _PLANAR_CONTRACTIONS = (:planarcontract!, :planaralloc_contract) function _insert_planar_operations(ex) if isexpr(ex, :call) @@ -99,6 +100,29 @@ function _insert_planar_operations(ex) return ex end +""" + canonicalizeplanarindices(ex, partitions) + +Replace the index tuples of every planar 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 to be canonicalized at runtime. +""" +function canonicalizeplanarindices(ex, partitions) + if isexpr(ex, :call) && length(ex.args) ≥ 7 && ex.args[1] isa GlobalRef && + ex.args[1].mod === TensorKit && ex.args[1].name ∈ _PLANAR_CONTRACTIONS + A, pA, B, pB, pAB = ex.args[3], ex.args[4], ex.args[5], ex.args[6], ex.args[7] + 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[6], args[7] = 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..0421b8298 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 _index_partitions!(partitions, ex) + 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 -> _index_partitions!(partitions, a), 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) diff --git a/test/tensors/planar.jl b/test/tensors/planar.jl index 1f476e867..8c11f3b96 100644 --- a/test/tensors/planar.jl +++ b/test/tensors/planar.jl @@ -192,6 +192,22 @@ end @test !isempty(alloc_trailing) 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 From 30993e64ffdfd634d8d64c1938c3a4da48eed10f Mon Sep 17 00:00:00 2001 From: lkdvos Date: Sun, 13 Sep 2026 09:38:27 -0400 Subject: [PATCH 3/6] Drop `planaralloc_contract` Now that `@planar` emits canonical index tuples, its allocations no longer need to be canonicalized at runtime and can go through `tensoralloc_contract` again. Canonicalizing before the planar operations are inserted also means that the contraction and its allocation still share a single argument layout. Co-Authored-By: Claude Opus 5 (1M context) --- docs/src/Changelog.md | 4 +-- .../tensoroperations.jl | 2 -- ext/TensorKitMooncakeExt/utility.jl | 1 - src/planar/macros.jl | 2 +- src/planar/planaroperations.jl | 19 +---------- src/planar/postprocessors.jl | 34 ++++++++----------- test/tensors/planar.jl | 10 ++---- 7 files changed, 21 insertions(+), 51 deletions(-) diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index dbc19829b..fbb216427 100644 --- a/docs/src/Changelog.md +++ b/docs/src/Changelog.md @@ -22,11 +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 internal helpers `TensorKit.planar_contract_indices` and `TensorKit.planaralloc_contract`. ([#531](https://github.com/QuantumKitHub/TensorKit.jl/pull/531)) +- `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()`; planar code should allocate through `TensorKit.planaralloc_contract`, which canonicalizes them. ([#531](https://github.com/QuantumKitHub/TensorKit.jl/pull/531)) +- `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/ext/TensorKitChainRulesCoreExt/tensoroperations.jl b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl index 5e695c35b..82bd9b578 100644 --- a/ext/TensorKitChainRulesCoreExt/tensoroperations.jl +++ b/ext/TensorKitChainRulesCoreExt/tensoroperations.jl @@ -198,8 +198,6 @@ function ChainRulesCore.rrule( return C′, pullback end -@non_differentiable TensorKit.planaralloc_contract(args...) - function ChainRulesCore.rrule(::typeof(TensorKit.scalar), t::AbstractTensorMap) val = scalar(t) function scalar_pullback(Δval) diff --git a/ext/TensorKitMooncakeExt/utility.jl b/ext/TensorKitMooncakeExt/utility.jl index 145c83e0b..1bedf566c 100644 --- a/ext/TensorKitMooncakeExt/utility.jl +++ b/ext/TensorKitMooncakeExt/utility.jl @@ -19,7 +19,6 @@ Mooncake.tangent_type(::Type{<:HomSpace}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorstructure), AbstractTensorMap, Int, Bool} @zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorcontract_structure), AbstractTensorMap, Index2Tuple, Bool, AbstractTensorMap, Index2Tuple, Bool, Index2Tuple} -@zero_derivative DefaultCtx Tuple{typeof(TensorKit.planaralloc_contract), Any} @zero_derivative DefaultCtx Tuple{typeof(TensorKit.planar_trace), TensorKit.FusionTreePair, Index2Tuple, Index2Tuple} @zero_derivative DefaultCtx Tuple{typeof(TensorKit.has_shared_permute), AbstractTensorMap, Index2Tuple} diff --git a/src/planar/macros.jl b/src/planar/macros.jl index aceb8dad8..c47f5d670 100644 --- a/src/planar/macros.jl +++ b/src/planar/macros.jl @@ -19,10 +19,10 @@ function planarparser(planarexpr, kwargs...) temporaries = Vector{Symbol}() push!(parser.postprocessors, ex -> _annotate_temporaries(ex, temporaries)) push!(parser.postprocessors, ex -> _free_temporaries(ex, temporaries)) - push!(parser.postprocessors, _insert_planar_operations) 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 push!(parser.preprocessors, _construct_braidingtensors) diff --git a/src/planar/planaroperations.jl b/src/planar/planaroperations.jl index 0ad133dcf..7ef2050a6 100644 --- a/src/planar/planaroperations.jl +++ b/src/planar/planaroperations.jl @@ -256,7 +256,7 @@ spaces, so all space computations should use them. `A` and `B` can be anything s `codomainind` and `domainind`, in particular `AbstractTensorMap`s, `HomSpace`s and [`IndexPartition`](@ref)s. -See also [`planarcontract!`](@ref) and [`planaralloc_contract`](@ref). +See also [`planarcontract!`](@ref). """ function planar_contract_indices( A, (oindA, cindA)::Index2Tuple, @@ -286,20 +286,3 @@ function planar_contract_indices( pAB′ = (TupleTools.getindices(remap, pAB[1]), TupleTools.getindices(remap, pAB[2])) return (oindA′, cindA′), (cindB′, oindB′), pAB′ end - -""" - planaralloc_contract(TC, A, pA, B, pB, pAB, [istemp, allocator]) - -Allocate the destination of `planarcontract!(C, A, pA, B, pB, pAB, α, β)`. - -The planar counterpart of `TensorOperations.tensoralloc_contract`: the index tuples are -canonicalized with [`planar_contract_indices`](@ref) first, such that the space computation -only involves valid intermediate spaces. -""" -function planaralloc_contract( - TC, A, pA::Index2Tuple, B, pB::Index2Tuple, pAB::Index2Tuple, - istemp::Val = Val(false), allocator = TO.DefaultAllocator() - ) - pA′, pB′, pAB′ = planar_contract_indices(A, pA, B, pB, pAB) - return TO.tensoralloc_contract(TC, A, pA′, false, B, pB′, false, pAB′, istemp, allocator) -end diff --git a/src/planar/postprocessors.jl b/src/planar/postprocessors.jl index de5a08c08..357d6c466 100644 --- a/src/planar/postprocessors.jl +++ b/src/planar/postprocessors.jl @@ -53,8 +53,10 @@ end # TODO: replace _planarmethod with planarmethod in everything below const _PLANAR_OPERATIONS = (:planaradd!, :planartrace!, :planarcontract!) -const _PLANAR_ALLOCATIONS = (:planaralloc_contract,) -const _PLANAR_CONTRACTIONS = (:planarcontract!, :planaralloc_contract) + +# `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) @@ -80,14 +82,6 @@ function _insert_planar_operations(ex) ex.head, GlobalRef(TensorKit, Symbol(:planartrace!)), map(_insert_planar_operations, ex.args[2:end])... ) - elseif ex.args[1] == GlobalRef(TensorOperations, :tensoralloc_contract) - conjB = popat!(ex.args, 8) - conjA = popat!(ex.args, 5) - @assert !conjA && !conjB "conj flags should be disabled ($conjA), ($conjB)" - return Expr( - ex.head, GlobalRef(TensorKit, Symbol(:planaralloc_contract)), - map(_insert_planar_operations, ex.args[2:end])... - ) elseif ex.args[1] in TensorOperations.tensoroperationsfunctions return Expr( ex.head, GlobalRef(TensorOperations, ex.args[1]), @@ -103,19 +97,20 @@ end """ canonicalizeplanarindices(ex, partitions) -Replace the index tuples of every planar 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 to be canonicalized at runtime. +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) ≥ 7 && ex.args[1] isa GlobalRef && - ex.args[1].mod === TensorKit && ex.args[1].name ∈ _PLANAR_CONTRACTIONS - A, pA, B, pB, pAB = ex.args[3], ex.args[4], ex.args[5], ex.args[6], ex.args[7] + 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[6], args[7] = planar_contract_indices(WA, pA, WB, pB, pAB) + args[4], args[7], args[9] = planar_contract_indices(WA, pA, WB, pB, pAB) return Expr(:call, args...) end end @@ -149,11 +144,10 @@ end """ insertplanarallocator(ex, allocator) -Insert the allocator argument into the tensor operation methods `planaradd!`, `planartrace!`, -`planarcontract!`, and `planaralloc_contract`. +Insert the allocator argument into the tensor operation methods `planaradd!`, `planartrace!`, and `planarcontract!`. See also: [`TensorOperations.insertallocator`](@ref). """ function insertplanarallocator(ex, allocator) - return _insertargument(ex, allocator, (_PLANAR_OPERATIONS..., _PLANAR_ALLOCATIONS...)) + return _insertargument(ex, allocator, _PLANAR_OPERATIONS) end diff --git a/test/tensors/planar.jl b/test/tensors/planar.jl index 8c11f3b96..72a933b21 100644 --- a/test/tensors/planar.jl +++ b/test/tensors/planar.jl @@ -161,6 +161,7 @@ end T = ComplexF64 @testset "backend and allocator insertion" begin + 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 @@ -174,22 +175,17 @@ end ex = @macroexpand @planar backend = MarkerBackend() C[i; j] := A[i; k l] * τ[k l; m n] * B[m n; j] - trailing = planartrailing(ex, (:planaradd!, :planartrace!, :planarcontract!)) + 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, (:planaradd!, :planartrace!, :planarcontract!, :planaralloc_contract) - ) + trailing = planartrailing(ex, PLANAR_OPERATIONS) @test !isempty(trailing) @test all(==(:(MarkerAllocator())), trailing) @test occursin("DefaultBackend", string(ex)) - - alloc_trailing = planartrailing(ex, (:planaralloc_contract,)) - @test !isempty(alloc_trailing) end @testset "canonical index tuples" begin From 988deb3d63f5f5b652815a6033e73b020265e602 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Thu, 17 Sep 2026 13:59:47 -0400 Subject: [PATCH 4/6] if -> elseif --- src/planar/planaroperations.jl | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/planar/planaroperations.jl b/src/planar/planaroperations.jl index 7ef2050a6..5866f7ef9 100644 --- a/src/planar/planaroperations.jl +++ b/src/planar/planaroperations.jl @@ -271,8 +271,7 @@ function planar_contract_indices( # 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))]) - end - if isempty(oindB′) && !isempty(cindB) + elseif isempty(oindB′) && !isempty(cindB) cindB′ = _rotate_to(cindB′, cindB[something(findfirst(==(first(cindA′)), cindA))]) end TupleTools.sort(tuple.(cindA′, cindB′)) == TupleTools.sort(tuple.(cindA, cindB)) || From 4494902a602c1448d5c02ed7e854e62297f6bd99 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Sun, 20 Sep 2026 10:38:25 -0400 Subject: [PATCH 5/6] Address review comments on planar preprocessors Rename `_index_partitions!` to `_record_index_partitions!` so the preprocessor name starts with a verb, and take `ex` as its first argument for consistency with the other preprocessors. Add the missing `!` to `_decompose_planar_contractions!`, which likewise mutates its accumulator. Also test the planar contraction with explicit parentheses, so that both association orders of the three-factor contraction are covered. Co-Authored-By: Claude Opus 5 (1M context) --- src/planar/macros.jl | 4 ++-- src/planar/postprocessors.jl | 2 +- src/planar/preprocessors.jl | 10 +++++----- test/tensors/contractions.jl | 3 +++ 4 files changed, 11 insertions(+), 8 deletions(-) diff --git a/src/planar/macros.jl b/src/planar/macros.jl index c47f5d670..a14695812 100644 --- a/src/planar/macros.jl +++ b/src/planar/macros.jl @@ -93,8 +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 -> _index_partitions!(partitions, ex)) + 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/postprocessors.jl b/src/planar/postprocessors.jl index 357d6c466..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) diff --git a/src/planar/preprocessors.jl b/src/planar/preprocessors.jl index 0421b8298..17a99eb73 100644 --- a/src/planar/preprocessors.jl +++ b/src/planar/preprocessors.jl @@ -78,13 +78,13 @@ 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 _index_partitions!(partitions, ex) +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 -> _index_partitions!(partitions, a), ex.args) + foreach(a -> _record_index_partitions!(a, partitions), ex.args) end return ex end @@ -462,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 @@ -484,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/test/tensors/contractions.jl b/test/tensors/contractions.jl index 195e7851e..1132809dc 100644 --- a/test/tensors/contractions.jl +++ b/test/tensors/contractions.jl @@ -75,6 +75,9 @@ for V in spacelist # 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 From bb637c8f41a41e1ce196b96fb711f411f8938e00 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Sun, 20 Sep 2026 10:38:25 -0400 Subject: [PATCH 6/6] Normalize `One`/`Zero` scalars before BLAS in `mul!` MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `mul!(tC, tA, tB, α, β)` forwarded `VectorInterface.One`/`Zero` verbatim to LinearAlgebra's block `mul!`. Ordinary GEMM tolerates this, but a contraction of a tensor with its own adjoint dispatches to the HERK path, where `herk_wrapper!` calls `isreal` on the scalars and throws a `MethodError`. Map them to `true`/`false` first, which LinearAlgebra still recognizes as exact one and zero. Co-Authored-By: Claude Opus 5 (1M context) --- src/tensors/linalg.jl | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/src/tensors/linalg.jl b/src/tensors/linalg.jl index f0b000b0e..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)