From 66678b656586e2997665f750f4141e4b4c3809d4 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 18:30:30 -0400 Subject: [PATCH 01/33] Open the v0.21 breaking round Accumulates the contract and matricize interface redesign. The version stays at 0.21.0-DEV until the release PR strips the suffix. Co-Authored-By: Claude Opus 5 (1M context) --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index 146f410b..9cf7ac47 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.20.1" +version = "0.21.0-DEV" authors = ["ITensor developers and contributors"] [workspace] From 4a63a19b1ebd2fcc1277caee89e0a4aa1dc442af Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 18:36:30 -0400 Subject: [PATCH 02/33] Delete the out-of-place unmatricizeperm It had no callers outside its own forwarding method, and its behaviour is covered by the in-place form. Also raises the subproject compat bounds the round-opening bump missed. Co-Authored-By: Claude Opus 5 (1M context) --- docs/Project.toml | 2 +- examples/Project.toml | 2 +- src/matricize.jl | 29 +---------------------------- test/Project.toml | 2 +- test/test_basics.jl | 14 +++----------- 5 files changed, 7 insertions(+), 42 deletions(-) diff --git a/docs/Project.toml b/docs/Project.toml index e02c5956..d9d7449e 100644 --- a/docs/Project.toml +++ b/docs/Project.toml @@ -11,4 +11,4 @@ path = ".." Documenter = "1.8.1" ITensorFormatter = "0.2.27" Literate = "2.20.1" -TensorAlgebra = "0.20" +TensorAlgebra = "0.21" diff --git a/examples/Project.toml b/examples/Project.toml index b7e5f73e..194d5b6d 100644 --- a/examples/Project.toml +++ b/examples/Project.toml @@ -5,4 +5,4 @@ TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" path = ".." [compat] -TensorAlgebra = "0.20" +TensorAlgebra = "0.21" diff --git a/src/matricize.jl b/src/matricize.jl index 972cb072..5ede4d82 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -226,7 +226,7 @@ end # domain groups, given codomain-facing (un-dualized), the same convention as `similar_map`. A # matricize style stores the domain axes dualized, so its overload re-dualizes them with `conj` # (a no-op on a dense axis). This is the primary overload point for new matricize styles. -# Permutation is handled separately by `unmatricizeperm`, so `unmatricize` never has to +# Permutation is handled separately by `unmatricizeperm!`, so `unmatricize` never has to # disambiguate axis tuples from permutation tuples regardless of how unconstrained `m` and the # axes are. function unmatricize(style::MatricizeStyle, m, axes_codomain, axes_domain) @@ -245,33 +245,6 @@ function bipartition_axes(t::Tuple, split...) return axes_codomain, conj.(axes_domain) end -# Inverse-bipermutation form: split `axes_dest` into codomain/domain groups reordered by the -# inverse bipermutation, unmatricize in that order, then permute back. -function unmatricizeperm( - m, axes_dest, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} - ) - return unmatricizeperm( - MatricizeStyle(m), - m, - axes_dest, - invperm_codomain, - invperm_domain - ) -end -function unmatricizeperm( - style::MatricizeStyle, m, axes_dest, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} - ) - invbiperm = BiTuple(invperm_codomain, invperm_domain) - length(axes_dest) == length(invbiperm) || - throw(ArgumentError("axes do not match permutation")) - axes_codomain, axes_domain = bipartition_axes(axes_dest, invbiperm) - a12 = unmatricize(style, m, axes_codomain, axes_domain) - biperm_dest = BiTuple(Tuple(invperm(invbiperm)), Val(length_codomain(invbiperm))) - return bipermutedims(a12, biperm_dest) -end - function unmatricizeperm!( a_dest, m, invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} diff --git a/test/Project.toml b/test/Project.toml index 08c7d265..0ba4cc34 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -37,7 +37,7 @@ Random = "1.10" SafeTestsets = "0.1" StableRNGs = "1.0.2" Suppressor = "0.2" -TensorAlgebra = "0.20" +TensorAlgebra = "0.21" TensorKit = "0.17" TensorOperations = "5.1.4" Test = "1.10" diff --git a/test/test_basics.jl b/test/test_basics.jl index 9e285f86..d5f7686f 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -3,7 +3,7 @@ using EllipsisNotation: var".." using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, contract!, contractadd!, length_codomain, length_domain, matricizeperm, unmatricize, - unmatricizeperm, unmatricizeperm! + unmatricizeperm! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -130,25 +130,18 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test eltype(a) === elt @test a ≈ a0 - a = unmatricizeperm(m, axes0, (1, 2), (3, 4)) - @test eltype(a) === elt - @test a ≈ a0 - perm_codomain = (4, 2) perm_domain = (1, 3) invperm_codomain = (3, 2) invperm_domain = (4, 1) perm = (4, 2, 1, 3) - a = unmatricizeperm(m, map(i -> axes0[i], perm), invperm_codomain, invperm_domain) - @test eltype(a) === elt - @test a ≈ permutedims(a0, perm) - a = similar(a0) unmatricizeperm!(a, m, (1, 2), (3, 4)) @test a ≈ a0 m1 = matricizeperm(a0, perm_codomain, perm_domain) - a = unmatricizeperm(m1, axes0, perm_codomain, perm_domain) + a = similar(a0) + unmatricizeperm!(a, m1, perm_codomain, perm_domain) @test a ≈ a0 a1 = permutedims(a0, perm) @@ -174,7 +167,6 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test a isa Array{elt, 0} @test a[] == m[1, 1] - @test_throws ArgumentError unmatricizeperm(m, (), (1, 2), (3,)) @test_throws ArgumentError unmatricizeperm!(m, m, (1, 2), (3,)) end From 17521648f7bd6de2ff40b21ed270d4e59f32eed0 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 18:44:13 -0400 Subject: [PATCH 03/33] Name the unmatricize bipermutation for what it is The argument maps the destination's dimension order to the matrix's; it is not intrinsically an inverse, so invperm_ described one caller's derivation rather than the parameter. Co-Authored-By: Claude Opus 5 (1M context) --- src/matricize.jl | 25 +++++++++++++++---------- test/test_matricizestyle.jl | 4 ++-- 2 files changed, 17 insertions(+), 12 deletions(-) diff --git a/src/matricize.jl b/src/matricize.jl index 5ede4d82..10465a58 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -215,10 +215,10 @@ end ismatricizeview(::MatricizeStyle, a, ndims_codomain::Val) = false function ismatricizeview( style::MatricizeStyle, a, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - isidentityperm((invperm_codomain..., invperm_domain...)) || return false - return ismatricizeview(style, a, Val(length(invperm_codomain))) + isidentityperm((perm_codomain..., perm_domain...)) || return false + return ismatricizeview(style, a, Val(length(perm_codomain))) end # ==================================== unmatricize ======================================= @@ -245,22 +245,27 @@ function bipartition_axes(t::Tuple, split...) return axes_codomain, conj.(axes_domain) end +# The bipermutation maps the destination's dimension order to the matrix's: `axes(a_dest)` grouped +# by it gives the legs in `m`'s order, and the result is permuted back by its inverse. It is not +# intrinsically an inverse permutation — the matricized-contraction destination path happens to +# derive it as `invperm(biperm_dest)`, while a `matricizeperm`/`unmatricizeperm!` round trip passes +# the same forward bipermutation to both. function unmatricizeperm!( a_dest, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - return unmatricizeperm!(MatricizeStyle(m), a_dest, m, invperm_codomain, invperm_domain) + return unmatricizeperm!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) end function unmatricizeperm!( style::MatricizeStyle, a_dest, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - invbiperm = BiTuple(invperm_codomain, invperm_domain) - ndims(a_dest) == length(invbiperm) || + biperm_src = BiTuple(perm_codomain, perm_domain) + ndims(a_dest) == length(biperm_src) || throw(ArgumentError("destination does not match permutation")) - axes_codomain, axes_domain = bipartition_axes(axes(a_dest), invbiperm) + axes_codomain, axes_domain = bipartition_axes(axes(a_dest), biperm_src) a_perm = unmatricize(style, m, axes_codomain, axes_domain) - biperm_dest = BiTuple(Tuple(invperm(invbiperm)), Val(length_codomain(invbiperm))) + biperm_dest = BiTuple(Tuple(invperm(biperm_src)), Val(length_codomain(biperm_src))) return bipermutedims!(a_dest, a_perm, biperm_dest) end diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index c43627f1..c7851e59 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -18,10 +18,10 @@ module MatricizeStyleTestUtils end function TA.unmatricizeperm!( ::MyArrayMatricize, a_dest::MyArray, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) TA.unmatricizeperm!( - TA.ReshapeMatricize(), a_dest.parent, m, invperm_codomain, invperm_domain + TA.ReshapeMatricize(), a_dest.parent, m, perm_codomain, perm_domain ) return a_dest end From 3a9139dc5244fa0e35a338eca4092de73a25066d Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 18:50:18 -0400 Subject: [PATCH 04/33] Merge unmatricizeperm! into unmatricize! The two shapes are disjoint at the trailing argument, a Val split spec against a pair of permutation tuples, so one name carries both and the perm marker stops earning its place. Co-Authored-By: Claude Opus 5 (1M context) --- src/contract/contract_matricize.jl | 4 ++-- src/matricize.jl | 17 +++++++++-------- test/test_basics.jl | 10 +++++----- test/test_matricizestyle.jl | 4 ++-- 4 files changed, 18 insertions(+), 17 deletions(-) diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 9e04f5f0..1593d734 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -34,13 +34,13 @@ function contractopadd!( # overwrites `a_dest` in full. a_dest_mat = a1_mat * a2_mat isone(α) || scale!(a_dest_mat, α) - unmatricizeperm!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) + unmatricize!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) else # `a_dest`'s data contributes through `β`, so gather it, multiply into the gathered # copy, and scatter back. a_dest_mat = matricizecopy(output_style, a_dest, invperm_codomain, invperm_domain) mul!(a_dest_mat, a1_mat, a2_mat, α, β) - unmatricizeperm!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) + unmatricize!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) end return a_dest end diff --git a/src/matricize.jl b/src/matricize.jl index 10465a58..b7e67115 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -226,7 +226,8 @@ end # domain groups, given codomain-facing (un-dualized), the same convention as `similar_map`. A # matricize style stores the domain axes dualized, so its overload re-dualizes them with `conj` # (a no-op on a dense axis). This is the primary overload point for new matricize styles. -# Permutation is handled separately by `unmatricizeperm!`, so `unmatricize` never has to +# Permutation is handled by the bipermutation form of `unmatricize!`, so out-of-place `unmatricize` +# never has to # disambiguate axis tuples from permutation tuples regardless of how unconstrained `m` and the # axes are. function unmatricize(style::MatricizeStyle, m, axes_codomain, axes_domain) @@ -248,15 +249,15 @@ end # The bipermutation maps the destination's dimension order to the matrix's: `axes(a_dest)` grouped # by it gives the legs in `m`'s order, and the result is permuted back by its inverse. It is not # intrinsically an inverse permutation — the matricized-contraction destination path happens to -# derive it as `invperm(biperm_dest)`, while a `matricizeperm`/`unmatricizeperm!` round trip passes +# derive it as `invperm(biperm_dest)`, while a `matricizeperm`/`unmatricize!` round trip passes # the same forward bipermutation to both. -function unmatricizeperm!( +function unmatricize!( a_dest, m, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - return unmatricizeperm!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) + return unmatricize!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) end -function unmatricizeperm!( +function unmatricize!( style::MatricizeStyle, a_dest, m, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) @@ -269,14 +270,14 @@ function unmatricizeperm!( return bipermutedims!(a_dest, a_perm, biperm_dest) end -# In-place split-axes counterpart of `unmatricize`, as `unmatricizeperm!` is of `unmatricizeperm`: +# In-place counterpart of `unmatricize`: # scatter the fused matrix `m` back into `a_dest`'s existing storage across the codomain/domain -# split at `ndims_codomain`. The split applies no permutation, so this is `unmatricizeperm!` at the +# split at `ndims_codomain`. The split applies no permutation, so this is the bipermutation form at the # trivial bipermutation, reusing its in-place block scatter (no intermediate `unmatricize` copy). function unmatricize!(style::MatricizeStyle, a_dest, m, ndims_codomain::Val) K = unval(ndims_codomain) N = ndims(a_dest) - return unmatricizeperm!( + return unmatricize!( style, a_dest, m, diff --git a/test/test_basics.jl b/test/test_basics.jl index d5f7686f..64d4eee0 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -3,7 +3,7 @@ using EllipsisNotation: var".." using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, contract!, contractadd!, length_codomain, length_domain, matricizeperm, unmatricize, - unmatricizeperm! + unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -136,17 +136,17 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int invperm_domain = (4, 1) perm = (4, 2, 1, 3) a = similar(a0) - unmatricizeperm!(a, m, (1, 2), (3, 4)) + unmatricize!(a, m, (1, 2), (3, 4)) @test a ≈ a0 m1 = matricizeperm(a0, perm_codomain, perm_domain) a = similar(a0) - unmatricizeperm!(a, m1, perm_codomain, perm_domain) + unmatricize!(a, m1, perm_codomain, perm_domain) @test a ≈ a0 a1 = permutedims(a0, perm) a = similar(a1) - unmatricizeperm!(a, m, invperm_codomain, invperm_domain) + unmatricize!(a, m, invperm_codomain, invperm_domain) @test a ≈ a1 a = unmatricize(reshape(a0, 1, 120), (), axes0) @@ -167,7 +167,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test a isa Array{elt, 0} @test a[] == m[1, 1] - @test_throws ArgumentError unmatricizeperm!(m, m, (1, 2), (3,)) + @test_throws ArgumentError unmatricize!(m, m, (1, 2), (3,)) end @testset "contract eltype widens like a matrix product" begin diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index c7851e59..038ac12b 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -16,11 +16,11 @@ module MatricizeStyleTestUtils function TA.matricizecopy(::MyArrayMatricize, a::MyArray, ndims_codomain::Val) return TA.matricizecopy(TA.ReshapeMatricize(), a.parent, ndims_codomain) end - function TA.unmatricizeperm!( + function TA.unmatricize!( ::MyArrayMatricize, a_dest::MyArray, m, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - TA.unmatricizeperm!( + TA.unmatricize!( TA.ReshapeMatricize(), a_dest.parent, m, perm_codomain, perm_domain ) return a_dest From 00da2dc4d40224a6ff44423110ef49bd6b5c029b Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 19:00:41 -0400 Subject: [PATCH 05/33] Report an unusable contraction keyword by name The entry points collect trailing keywords and forward them to the resolver, whose methods declared none, so any unrecognized keyword surfaced as a MethodError on an internal function. Co-Authored-By: Claude Opus 5 (1M context) --- src/contract/contractalgorithm.jl | 35 ++++++++++++++++++++++++------- test/test_basics.jl | 18 ++++++++++++++++ 2 files changed, 46 insertions(+), 7 deletions(-) diff --git a/src/contract/contractalgorithm.jl b/src/contract/contractalgorithm.jl index 02fb0183..c172ccff 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -23,17 +23,38 @@ Base.@kwdef struct TensorOperationsAlgorithm{Backend, Allocator} <: ContractAlgo allocator::Allocator = nothing end -function select_contract_algorithm(algorithm, a1, a2) - return error("Not implemented.") +# The contraction entry points collect trailing keywords and forward them here, so these accept +# `kwargs...` even though no `ContractAlgorithm` is configurable by keyword yet. Without it an +# unrecognized keyword surfaces as a `MethodError` on this internal function rather than as a +# complaint about the keyword the caller actually passed. +function reject_algorithm_kwargs(algorithm; kwargs...) + isempty(kwargs) && return nothing + names = join(map(k -> "`$k`", collect(keys(kwargs))), ", ") + return throw( + ArgumentError( + "unsupported keyword argument(s) $names for contraction algorithm `$(nameof(typeof(algorithm)))`" + ) + ) end -function select_contract_algorithm(algorithm::ContractAlgorithm, a1, a2) + +function select_contract_algorithm(algorithm, a1, a2; kwargs...) + return throw( + ArgumentError( + "`$algorithm` is not a contraction algorithm; pass a `ContractAlgorithm` as `alg`" + ) + ) +end +function select_contract_algorithm(algorithm::ContractAlgorithm, a1, a2; kwargs...) + reject_algorithm_kwargs(algorithm; kwargs...) return algorithm end -function select_contract_algorithm(algorithm::DefaultContractAlgorithm, a1, a2) - return default_contract_algorithm(a1, a2) +function select_contract_algorithm(algorithm::DefaultContractAlgorithm, a1, a2; kwargs...) + return default_contract_algorithm(a1, a2; kwargs...) end -function default_contract_algorithm(a1, a2) - return default_contract_algorithm(typeof(a1), typeof(a2)) +function default_contract_algorithm(a1, a2; kwargs...) + algorithm = default_contract_algorithm(typeof(a1), typeof(a2)) + reject_algorithm_kwargs(algorithm; kwargs...) + return algorithm end function default_contract_algorithm(A1::Type{<:AbstractArray}, A2::Type{<:AbstractArray}) return Matricize(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) diff --git a/test/test_basics.jl b/test/test_basics.jl index 64d4eee0..0d789da2 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -170,6 +170,24 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test_throws ArgumentError unmatricize!(m, m, (1, 2), (3,)) end + @testset "contraction algorithm selection rejects unusable keywords" begin + a1 = randn(2, 3) + a2 = randn(3, 4) + # A keyword no algorithm can consume must name itself, not surface as a `MethodError` + # from inside the resolver. + @test_throws ArgumentError contract((1, 3), a1, (1, 2), a2, (2, 3); nonsense = 1) + @test_throws ArgumentError contract( + (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize(), nonsense = 1 + ) + # A non-algorithm passed as `alg` says so rather than erroring with "Not implemented". + @test_throws ArgumentError TensorAlgebra.select_contract_algorithm(:nope, a1, a2) + # The supported spellings still work. + @test contract((1, 3), a1, (1, 2), a2, (2, 3)) ≈ a1 * a2 + @test contract( + (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize() + ) ≈ a1 * a2 + end + @testset "contract eltype widens like a matrix product" begin a1 = ones(Bool, (2, 2)) a2 = ones(Bool, (2, 2)) From 5df5d28a982d64e24d1bc841578a82b5a13a6be1 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 19:28:50 -0400 Subject: [PATCH 06/33] Drop the unused Ellipsis bipermutation spelling Nothing passed `..`, and supporting it cost a dependency plus a typed/untyped method tier whose only job was normalizing it. Names the joint predicate `isbiperm` and routes the three duplicate validations through `check_biperm`. Co-Authored-By: Claude Opus 5 (1M context) --- Project.toml | 2 -- src/bituple.jl | 5 ++++ src/contract/allocate_output.jl | 2 +- src/matricize.jl | 48 +++------------------------------ test/Project.toml | 2 -- test/test_basics.jl | 1 - 6 files changed, 10 insertions(+), 50 deletions(-) diff --git a/Project.toml b/Project.toml index 9cf7ac47..30d47b0b 100644 --- a/Project.toml +++ b/Project.toml @@ -7,7 +7,6 @@ authors = ["ITensor developers and contributors"] projects = ["benchmark", "dev", "docs", "examples", "test"] [deps] -EllipsisNotation = "da5c29d0-fa7d-589e-88eb-ea29b0a81949" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" MatrixAlgebraKit = "6c742aac-3347-4629-af66-fc926824e5e4" Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c" @@ -29,7 +28,6 @@ TensorAlgebraTensorKitSectorsExt = "TensorKitSectors" TensorAlgebraTensorOperationsExt = "TensorOperations" [compat] -EllipsisNotation = "1.8" LinearAlgebra = "1.10" MatrixAlgebraKit = "0.6" Mooncake = "0.4.202, 0.5" diff --git a/src/bituple.jl b/src/bituple.jl index 17acf4cb..0501193e 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -70,3 +70,8 @@ function bipartition(t::Tuple, group1::Tuple, group2::Tuple) end # Split `t` by the two groups of a `BiTuple`. bipartition(t::Tuple, bt::BiTuple) = bipartition(t, bt.t1, bt.t2) + +# Whether `perm_codomain` and `perm_domain` together permute `1:n`, i.e. whether they are a valid +# bipartitioned permutation. The two halves only make sense jointly, so this takes them as a pair +# rather than leaving every caller to splat and call `isperm`. +isbiperm(perm_codomain, perm_domain) = isperm((perm_codomain..., perm_domain...)) diff --git a/src/contract/allocate_output.jl b/src/contract/allocate_output.jl index c09f6716..13b392c0 100644 --- a/src/contract/allocate_output.jl +++ b/src/contract/allocate_output.jl @@ -1,7 +1,7 @@ function check_biperm(a, perm_codomain, perm_domain) ndims(a) == length(perm_codomain) + length(perm_domain) || throw(ArgumentError("Invalid bipartitioned permutation")) - isperm((perm_codomain..., perm_domain...)) || + isbiperm(perm_codomain, perm_domain) || throw(ArgumentError("Invalid bipartitioned permutation")) return nothing end diff --git a/src/matricize.jl b/src/matricize.jl index b7e67115..4b88f0cb 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -1,4 +1,3 @@ -using EllipsisNotation: Ellipsis using LinearAlgebra: Diagonal # ===================================== MatricizeStyle ====================================== @@ -38,9 +37,7 @@ Non-mutating version of `bipermutedimsopadd!`: returns function permutedimsop(op, src, perm_codomain, perm_domain) # Validate against `src` here: `bipermutedimsopadd!`'s `check_input` compares against `dest`, # which `allocate_output` builds from the same perms, so it cannot catch a non-covering perm. - perm = (perm_codomain..., perm_domain...) - (ndims(src) == length(perm) && isperm(perm)) || - throw(ArgumentError("Invalid bipermutation")) + check_biperm(src, perm_codomain, perm_domain) dest = allocate_output(permutedimsop, op, src, perm_codomain, perm_domain) return bipermutedimsopadd!(dest, op, src, perm_codomain, perm_domain, true, false) end @@ -132,39 +129,6 @@ function matricizeperm( return matricizeopperm(style, identity, a, perm_codomain, perm_domain) end -# Process inputs such as `EllipsisNotation.Ellipsis`. -function to_permblocks(a, permblocks::NTuple{2, Tuple{Vararg{Int}}}) - isperm((permblocks[1]..., permblocks[2]...)) || - throw(ArgumentError("Invalid bipermutation")) - return permblocks -end -# Like `setcomplement` is like `setdiff` but assumes t2 ⊆ t1. -function tuplesetcomplement(t1::NTuple{N1}, t2::NTuple{N2}) where {N1, N2} - t2 ⊆ t1 || throw(ArgumentError("t2 must be a subset of t1")) - return NTuple{N1 - N2}(setdiff(t1, t2)) -end -function to_permblocks( - a, permblocks::Tuple{Tuple{Ellipsis}, Tuple{Vararg{Int}}} - ) - permblocks1 = tuplesetcomplement(ntuple(identity, ndims(a)), permblocks[2]) - return (permblocks1, permblocks[2]) -end -function to_permblocks( - a, permblocks::Tuple{Tuple{Vararg{Int}}, Tuple{Ellipsis}} - ) - permblocks2 = tuplesetcomplement(ntuple(identity, ndims(a)), permblocks[1]) - return (permblocks[1], permblocks2) -end - -function matricizeperm(a, perm_codomain, perm_domain) - return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) -end -function matricizeperm( - style::MatricizeStyle, a, perm_codomain, perm_domain - ) - return matricizeperm(style, a, to_permblocks(a, (perm_codomain, perm_domain))...) -end - # ================================== matricizeopperm ===================================== """ @@ -177,13 +141,10 @@ Has "maybe alias" semantics: the result may be a view/wrapper aliasing `a` or a copy, depending on the matricize style and array type. The caller should treat the result as read-only. """ -function matricizeopperm(op, a, perm_codomain, perm_domain) - return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) -end function matricizeopperm( - style::MatricizeStyle, op, a, perm_codomain, perm_domain + op, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - return matricizeopperm(style, op, a, to_permblocks(a, (perm_codomain, perm_domain))...) + return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end # Whether `perm` is the identity permutation `(1, …, n)`. isidentityperm(perm::Tuple{Vararg{Int}}) = perm == ntuple(identity, length(perm)) @@ -196,8 +157,7 @@ function matricizeopperm( style::MatricizeStyle, op, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - ndims(a) == length(perm_codomain) + length(perm_domain) || - throw(ArgumentError("Invalid bipermutation")) + check_biperm(a, perm_codomain, perm_domain) op === identity && isidentityperm((perm_codomain..., perm_domain...)) && return matricize(style, a, Val(length(perm_codomain))) a_perm_op = permutedimsop(op, a, perm_codomain, perm_domain) diff --git a/test/Project.toml b/test/Project.toml index 0ba4cc34..a2924b93 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -3,7 +3,6 @@ Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e" Aqua = "4c88cf16-eb10-579e-8560-4a9242c79595" BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" BlockArrays = "8e7c35d0-a365-5155-bbbb-fb81a777f24e" -EllipsisNotation = "da5c29d0-fa7d-589e-88eb-ea29b0a81949" ITensorPkgSkeleton = "3d388ab1-018a-49f4-ae50-18094d5f71ea" JLArrays = "27aeb0d3-9eb9-45fb-866b-73c2ecf80fcb" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" @@ -27,7 +26,6 @@ Adapt = "4" Aqua = "0.8.9" BenchmarkTools = "1" BlockArrays = "1.6.1" -EllipsisNotation = "1.8" ITensorPkgSkeleton = "0.3.42" JLArrays = "0.3" LinearAlgebra = "<0.0.1, 1" diff --git a/test/test_basics.jl b/test/test_basics.jl index 0d789da2..218ce0fa 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,5 +1,4 @@ import TensorAlgebra -using EllipsisNotation: var".." using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, contract!, contractadd!, length_codomain, length_domain, matricizeperm, unmatricize, From 1ec949ffdeda1c67d4f4b61cb591609e6f7756e7 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 19:34:32 -0400 Subject: [PATCH 07/33] Leave the bipermutation arguments untyped A hook dispatches on the array or the style, never on the permutation, so annotating it narrowed what a backend may pass without buying any dispatch. Base leaves `permutedims`' perm untyped and validates at runtime. Co-Authored-By: Claude Opus 5 (1M context) --- src/factorizations.jl | 68 ++++++++++++++++++------------------- src/matricize.jl | 16 ++++----- src/matrixfunctions.jl | 4 +-- test/test_matricizestyle.jl | 2 +- 4 files changed, 45 insertions(+), 45 deletions(-) diff --git a/src/factorizations.jl b/src/factorizations.jl index 1ab39631..5f7903eb 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -29,7 +29,7 @@ for f in ( @eval begin function $f( style::MatricizeStyle, A, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) ndims(A) == length(perm_codomain) + length(perm_domain) || @@ -64,7 +64,7 @@ for f in ( @eval begin function $f( style::MatricizeStyle, A, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) A_mat = matricizeperm(style, A, perm_codomain, perm_domain) @@ -102,7 +102,7 @@ for f in ( end function $f( A, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) return $f(MatricizeStyle(A), A, perm_codomain, perm_domain; kwargs...) @@ -143,7 +143,7 @@ end """ TensorAlgebra.tr(A, labels_A, labels_codomain, labels_domain) - TensorAlgebra.tr(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) + TensorAlgebra.tr(A, perm_codomain, perm_domain) TensorAlgebra.tr(A, ndims_codomain::Val) Trace of a generic N-dimensional array `A` interpreted as a linear map from its domain to its @@ -173,7 +173,7 @@ end function tr(A, ndims_codomain::Val) return tr(MatricizeStyle(A), A, ndims_codomain) end -function tr(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) +function tr(A, perm_codomain, perm_domain) return LinearAlgebra.tr(matricizeperm(A, perm_codomain, perm_domain)) end function tr(A, labels_A, labels_codomain, labels_domain) @@ -184,7 +184,7 @@ end """ qr_compact(A, labels_A, labels_codomain, labels_domain; kwargs...) -> Q, R - qr_compact(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> Q, R + qr_compact(A, perm_codomain, perm_domain; kwargs...) -> Q, R qr_compact(A, ndims_codomain::Val; kwargs...) -> Q, R Compute the compact QR decomposition of a generic N-dimensional array, by interpreting it @@ -202,7 +202,7 @@ qr_compact """ qr_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> Q, R - qr_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> Q, R + qr_full(A, perm_codomain, perm_domain; kwargs...) -> Q, R qr_full(A, ndims_codomain::Val; kwargs...) -> Q, R Compute the full QR decomposition of a generic N-dimensional array, by interpreting it as @@ -220,7 +220,7 @@ qr_full """ lq_compact(A, labels_A, labels_codomain, labels_domain; kwargs...) -> L, Q - lq_compact(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> L, Q + lq_compact(A, perm_codomain, perm_domain; kwargs...) -> L, Q lq_compact(A, ndims_codomain::Val; kwargs...) -> L, Q Compute the compact LQ decomposition of a generic N-dimensional array, by interpreting it @@ -238,7 +238,7 @@ lq_compact """ lq_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> L, Q - lq_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> L, Q + lq_full(A, perm_codomain, perm_domain; kwargs...) -> L, Q lq_full(A, ndims_codomain::Val; kwargs...) -> L, Q Compute the full LQ decomposition of a generic N-dimensional array, by interpreting it as @@ -256,7 +256,7 @@ lq_full """ left_polar(A, labels_A, labels_codomain, labels_domain; kwargs...) -> W, P - left_polar(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> W, P + left_polar(A, perm_codomain, perm_domain; kwargs...) -> W, P left_polar(A, ndims_codomain::Val; kwargs...) -> W, P Compute the left polar decomposition of a generic N-dimensional array, by interpreting it as @@ -273,7 +273,7 @@ left_polar """ right_polar(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P, W - right_polar(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> P, W + right_polar(A, perm_codomain, perm_domain; kwargs...) -> P, W right_polar(A, ndims_codomain::Val; kwargs...) -> P, W Compute the right polar decomposition of a generic N-dimensional array, by interpreting it as @@ -290,7 +290,7 @@ right_polar """ left_orth(A, labels_A, labels_codomain, labels_domain; kwargs...) -> V, C - left_orth(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> V, C + left_orth(A, perm_codomain, perm_domain; kwargs...) -> V, C left_orth(A, ndims_codomain::Val; kwargs...) -> V, C Compute the left orthogonal decomposition of a generic N-dimensional array, by interpreting it as @@ -307,7 +307,7 @@ left_orth """ right_orth(A, labels_A, labels_codomain, labels_domain; kwargs...) -> C, V - right_orth(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> C, V + right_orth(A, perm_codomain, perm_domain; kwargs...) -> C, V right_orth(A, ndims_codomain::Val; kwargs...) -> C, V Compute the right orthogonal decomposition of a generic N-dimensional array, by interpreting it as @@ -383,7 +383,7 @@ end """ svd_compact(A, labels_A, labels_codomain, labels_domain; kwargs...) -> U, S, Vᴴ - svd_compact(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> U, S, Vᴴ + svd_compact(A, perm_codomain, perm_domain; kwargs...) -> U, S, Vᴴ svd_compact(A, ndims_codomain::Val; kwargs...) -> U, S, Vᴴ Compute the compact (thin) SVD of a generic N-dimensional array, by interpreting it as a @@ -396,7 +396,7 @@ svd_compact """ svd_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> U, S, Vᴴ - svd_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> U, S, Vᴴ + svd_full(A, perm_codomain, perm_domain; kwargs...) -> U, S, Vᴴ svd_full(A, ndims_codomain::Val; kwargs...) -> U, S, Vᴴ Compute the full (thick) SVD of a generic N-dimensional array, by interpreting it as a @@ -409,7 +409,7 @@ svd_full """ svd_trunc(A, labels_A, labels_codomain, labels_domain; trunc, kwargs...) -> U, S, Vᴴ, ϵ - svd_trunc(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; trunc, kwargs...) -> U, S, Vᴴ, ϵ + svd_trunc(A, perm_codomain, perm_domain; trunc, kwargs...) -> U, S, Vᴴ, ϵ svd_trunc(A, ndims_codomain::Val; trunc, kwargs...) -> U, S, Vᴴ, ϵ Compute the truncated SVD of a generic N-dimensional array, by interpreting it as a linear @@ -446,7 +446,7 @@ svd_trunc """ svd_vals(A, labels_A, labels_codomain, labels_domain) -> S - svd_vals(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) -> S + svd_vals(A, perm_codomain, perm_domain) -> S svd_vals(A, ndims_codomain::Val) -> S Compute the singular values of a generic N-dimensional array, by interpreting it as a @@ -459,7 +459,7 @@ svd_vals """ eigh_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> D, V - eigh_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> D, V + eigh_full(A, perm_codomain, perm_domain; kwargs...) -> D, V eigh_full(A, ndims_codomain::Val; kwargs...) -> D, V Compute the eigenvalue decomposition of a generic N-dimensional array interpreted as a @@ -472,7 +472,7 @@ eigh_full """ eig_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> D, V - eig_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> D, V + eig_full(A, perm_codomain, perm_domain; kwargs...) -> D, V eig_full(A, ndims_codomain::Val; kwargs...) -> D, V Compute the eigenvalue decomposition of a generic N-dimensional array interpreted as a @@ -486,7 +486,7 @@ eig_full """ eigh_trunc(A, labels_A, labels_codomain, labels_domain; trunc, kwargs...) -> D, V - eigh_trunc(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; trunc, kwargs...) -> D, V + eigh_trunc(A, perm_codomain, perm_domain; trunc, kwargs...) -> D, V eigh_trunc(A, ndims_codomain::Val; trunc, kwargs...) -> D, V Truncated Hermitian eigenvalue decomposition, like [`eigh_full`](@ref) but keeping only the @@ -498,7 +498,7 @@ eigh_trunc """ eig_trunc(A, labels_A, labels_codomain, labels_domain; trunc, kwargs...) -> D, V - eig_trunc(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; trunc, kwargs...) -> D, V + eig_trunc(A, perm_codomain, perm_domain; trunc, kwargs...) -> D, V eig_trunc(A, ndims_codomain::Val; trunc, kwargs...) -> D, V Truncated general eigenvalue decomposition, like [`eig_full`](@ref) but keeping only the @@ -510,7 +510,7 @@ eig_trunc """ eigh_vals(A, labels_A, labels_codomain, labels_domain; kwargs...) -> D - eigh_vals(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> D + eigh_vals(A, perm_codomain, perm_domain; kwargs...) -> D eigh_vals(A, ndims_codomain::Val; kwargs...) -> D Compute the eigenvalues of a generic N-dimensional array interpreted as a Hermitian linear @@ -522,7 +522,7 @@ eigh_vals """ eig_vals(A, labels_A, labels_codomain, labels_domain; kwargs...) -> D - eig_vals(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> D + eig_vals(A, perm_codomain, perm_domain; kwargs...) -> D eig_vals(A, ndims_codomain::Val; kwargs...) -> D Compute the eigenvalues of a generic N-dimensional array interpreted as a general @@ -535,7 +535,7 @@ eig_vals """ left_null(A, labels_A, labels_codomain, labels_domain; kwargs...) -> N - left_null(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> N + left_null(A, perm_codomain, perm_domain; kwargs...) -> N left_null(A, ndims_codomain::Val; kwargs...) -> N Compute the left nullspace of a generic N-dimensional array, by interpreting it as @@ -572,7 +572,7 @@ end """ right_null(A, labels_A, labels_codomain, labels_domain; kwargs...) -> Nᴴ - right_null(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> Nᴴ + right_null(A, perm_codomain, perm_domain; kwargs...) -> Nᴴ right_null(A, ndims_codomain::Val::Val; kwargs...) -> Nᴴ Compute the right nullspace of a generic N-dimensional array, by interpreting it as @@ -609,7 +609,7 @@ end """ gram_eigh_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> X - gram_eigh_full(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> X + gram_eigh_full(A, perm_codomain, perm_domain; kwargs...) -> X gram_eigh_full(A, ndims_codomain::Val; kwargs...) -> X Gram factorization of a generic N-dimensional array, interpreting it as a @@ -665,7 +665,7 @@ end """ gram_eigh_full_with_pinv(A, labels_A, labels_codomain, labels_domain; kwargs...) -> X, Y - gram_eigh_full_with_pinv(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> X, Y + gram_eigh_full_with_pinv(A, perm_codomain, perm_domain; kwargs...) -> X, Y gram_eigh_full_with_pinv(A, ndims_codomain::Val; kwargs...) -> X, Y Like [`gram_eigh_full`](@ref), but additionally returns `Y ≈ pinv(X)` such @@ -727,7 +727,7 @@ end """ sqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P - sqrth_safe(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> P + sqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P sqrth_safe(A, ndims_codomain::Val; kwargs...) -> P Square root of a generic N-dimensional array, interpreting it as a @@ -750,7 +750,7 @@ sqrth_safe """ invsqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P - invsqrth_safe(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> P + invsqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P invsqrth_safe(A, ndims_codomain::Val; kwargs...) -> P Pseudo-inverse square root of a generic N-dimensional array, interpreting @@ -784,7 +784,7 @@ end """ project_hermitian(A, labels_A, labels_codomain, labels_domain; kwargs...) -> H - project_hermitian(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> H + project_hermitian(A, perm_codomain, perm_domain; kwargs...) -> H project_hermitian(A, ndims_codomain::Val; kwargs...) -> H Hermitian part `(M + M') / 2` of a generic N-dimensional array, interpreting @@ -804,7 +804,7 @@ end """ sqrth_invsqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P, Pinv - sqrth_invsqrth_safe(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs...) -> P, Pinv + sqrth_invsqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P, Pinv sqrth_invsqrth_safe(A, ndims_codomain::Val; kwargs...) -> P, Pinv Square root and pseudo-inverse square root of a generic N-dimensional @@ -833,7 +833,7 @@ end """ TensorAlgebra.one(A, labels_A, labels_codomain, labels_domain) -> Id - TensorAlgebra.one(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) -> Id + TensorAlgebra.one(A, perm_codomain, perm_domain) -> Id TensorAlgebra.one(A, ndims_codomain::Val) -> Id Construct the identity operator tensor whose shape mirrors `A`, interpreted as a @@ -901,14 +901,14 @@ end # `TensorMap` backend fills the identity through TensorKit rather than MatrixAlgebraKit). function one( style::MatricizeStyle, A, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) A_perm = bipermutedims(A, perm_codomain, perm_domain) return one!!(style, A_perm, Val(length(perm_codomain)); kwargs...) end function one( - A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs... + A, perm_codomain, perm_domain; kwargs... ) return one(MatricizeStyle(A), A, perm_codomain, perm_domain; kwargs...) end diff --git a/src/matricize.jl b/src/matricize.jl index 4b88f0cb..ac146d93 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -107,7 +107,7 @@ end # guaranteed to be a copy. function matricizecopy( style::MatricizeStyle, a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) a_perm = bipermutedims(a, perm_codomain, perm_domain) return matricize(style, a_perm, Val(length(perm_codomain))) @@ -115,7 +115,7 @@ end function matricizeperm( a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) end @@ -124,7 +124,7 @@ end # `matricizeopperm`. function matricizeperm( style::MatricizeStyle, a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) return matricizeopperm(style, identity, a, perm_codomain, perm_domain) end @@ -142,7 +142,7 @@ copy, depending on the matricize style and array type. The caller should treat t as read-only. """ function matricizeopperm( - op, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + op, a, perm_codomain, perm_domain ) return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end @@ -155,7 +155,7 @@ isidentityperm(perm::Tuple{Vararg{Int}}) = perm == ntuple(identity, length(perm) # alias `a` and must be treated as read-only, matching the docstring. function matricizeopperm( style::MatricizeStyle, op, a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) check_biperm(a, perm_codomain, perm_domain) op === identity && isidentityperm((perm_codomain..., perm_domain...)) && @@ -175,7 +175,7 @@ end ismatricizeview(::MatricizeStyle, a, ndims_codomain::Val) = false function ismatricizeview( style::MatricizeStyle, a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) isidentityperm((perm_codomain..., perm_domain...)) || return false return ismatricizeview(style, a, Val(length(perm_codomain))) @@ -213,13 +213,13 @@ end # the same forward bipermutation to both. function unmatricize!( a_dest, m, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) return unmatricize!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) end function unmatricize!( style::MatricizeStyle, a_dest, m, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) biperm_src = BiTuple(perm_codomain, perm_domain) ndims(a_dest) == length(biperm_src) || diff --git a/src/matrixfunctions.jl b/src/matrixfunctions.jl index ebbda116..1747244f 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -50,7 +50,7 @@ for f in MATRIX_FUNCTIONS function $f( style::MatricizeStyle, a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) a_mat = matricizeperm(style, a, perm_codomain, perm_domain) @@ -63,7 +63,7 @@ for f in MATRIX_FUNCTIONS end function $f( a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; + perm_codomain, perm_domain; kwargs... ) return $f(MatricizeStyle(a), a, perm_codomain, perm_domain; kwargs...) diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index 038ac12b..8d13c40b 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -18,7 +18,7 @@ module MatricizeStyleTestUtils end function TA.unmatricize!( ::MyArrayMatricize, a_dest::MyArray, m, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) TA.unmatricize!( TA.ReshapeMatricize(), a_dest.parent, m, perm_codomain, perm_domain From f5d7f6ad4eefd6a4463b7f8eb67299676d7b608a Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 19:36:35 -0400 Subject: [PATCH 08/33] Ask about identity bipermutations with both halves Both callers had a bipermutation in hand and were splatting it to reach `isidentityperm`, the same shape problem `isbiperm` fixed. Also drops the Ellipsis spellings from the matricize tests, which the removed normalizing tier supported. Co-Authored-By: Claude Opus 5 (1M context) --- src/bituple.jl | 8 ++++++++ src/matricize.jl | 6 ++---- test/test_basics.jl | 10 ++++------ 3 files changed, 14 insertions(+), 10 deletions(-) diff --git a/src/bituple.jl b/src/bituple.jl index 0501193e..f898ee75 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -75,3 +75,11 @@ bipartition(t::Tuple, bt::BiTuple) = bipartition(t, bt.t1, bt.t2) # bipartitioned permutation. The two halves only make sense jointly, so this takes them as a pair # rather than leaving every caller to splat and call `isperm`. isbiperm(perm_codomain, perm_domain) = isperm((perm_codomain..., perm_domain...)) + +# Whether `perm_codomain` and `perm_domain` are the identity bipermutation, i.e. leave every +# dimension where it is. Takes the halves for the same reason `isbiperm` does: both call sites had +# a bipermutation in hand and were splatting it to ask. +function isidentitybiperm(perm_codomain, perm_domain) + perm = (perm_codomain..., perm_domain...) + return perm == ntuple(identity, length(perm)) +end diff --git a/src/matricize.jl b/src/matricize.jl index ac146d93..70b2bebb 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -146,8 +146,6 @@ function matricizeopperm( ) return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end -# Whether `perm` is the identity permutation `(1, …, n)`. -isidentityperm(perm::Tuple{Vararg{Int}}) = perm == ntuple(identity, length(perm)) # The identity bipermutation is a no-op permute, so `matricize` runs directly on `a` (a view # for dense, a gather without the extra permute copy for graded); the fast path requires @@ -158,7 +156,7 @@ function matricizeopperm( perm_codomain, perm_domain ) check_biperm(a, perm_codomain, perm_domain) - op === identity && isidentityperm((perm_codomain..., perm_domain...)) && + op === identity && isidentitybiperm(perm_codomain, perm_domain) && return matricize(style, a, Val(length(perm_codomain))) a_perm_op = permutedimsop(op, a, perm_codomain, perm_domain) return matricize(style, a_perm_op, Val(length(perm_codomain))) @@ -177,7 +175,7 @@ function ismatricizeview( style::MatricizeStyle, a, perm_codomain, perm_domain ) - isidentityperm((perm_codomain..., perm_domain...)) || return false + isidentitybiperm(perm_codomain, perm_domain) || return false return ismatricizeview(style, a, Val(length(perm_codomain))) end diff --git a/test/test_basics.jl b/test/test_basics.jl index 218ce0fa..8f45811d 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -62,17 +62,15 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int a_fused = matricizeperm(a, (3, 1, 2), (4,)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(permutedims(a, (3, 1, 2, 4)), (24, 5)) - a_fused = matricizeperm(a, (..,), (3, 1)) + a_fused = matricizeperm(a, (2, 4), (3, 1)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(permutedims(a, (2, 4, 3, 1)), (15, 8)) - a_fused = matricizeperm(a, (3, 1), (..,)) - @test eltype(a_fused) === elt - @test a_fused ≈ reshape(permutedims(a, (3, 1, 2, 4)), (8, 15)) - a_fused = matricizeperm(a, (), (..,)) + # Degenerate splits: everything in the domain, then everything in the codomain. + a_fused = matricizeperm(a, (), (1, 2, 3, 4)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, (1, 120)) - a_fused = matricizeperm(a, (..,), ()) + a_fused = matricizeperm(a, (1, 2, 3, 4), ()) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, (120, 1)) From 77cc3c7a7f2fed42757b7deb3d5081b0ce162220 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 20:50:58 -0400 Subject: [PATCH 09/33] Give matricize one primitive per operation A style now implements four bipermutation hooks, `allocate_output`, `matricizeop!`, `matricizeopview` and `is_output_view`, and the copy and maybe-alias forms are derived. Allocation is a hook because only the style knows its fused axes, and because it is what makes the copy path terminate. Co-Authored-By: Claude Opus 5 (1M context) --- src/TensorAlgebra.jl | 2 +- src/contract/contract_matricize.jl | 16 ++- src/diagonal.jl | 6 +- src/factorizations.jl | 32 +++-- src/matricize.jl | 205 +++++++++++++++-------------- src/matrixfunctions.jl | 4 +- 6 files changed, 147 insertions(+), 118 deletions(-) diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 2853d1f3..0ac1bc95 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -9,7 +9,7 @@ export contract, contract!, dual, eig_full, eig_trunc, eig_vals, eigh_full, eigh if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public biperm, bipartition, cat_similar, concatenate, concatenate!, ContractAlgorithm, contractopadd!, data, datatype, directsum, flattenlinear, label_type, matricizeopperm, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" + "public biperm, bipartition, cat_similar, concatenate, concatenate!, ContractAlgorithm, contractopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" ) ) end diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 1593d734..0b5fc951 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -16,16 +16,20 @@ function contractopadd!( a1, biperm1_codomain, biperm1_domain, a2, biperm2_codomain, biperm2_domain ) - a1_mat = matricizeopperm( + a1_mat = matricizeop( algorithm.left_matricize_style, op1, a1, biperm1_codomain, biperm1_domain ) - a2_mat = matricizeopperm( + a2_mat = matricizeop( algorithm.right_matricize_style, op2, a2, biperm2_codomain, biperm2_domain ) output_style = algorithm.output_matricize_style - if ismatricizeview(output_style, a_dest, invperm_codomain, invperm_domain) + if is_output_view( + matricizeop, output_style, identity, a_dest, invperm_codomain, invperm_domain + ) # The matricization shares `a_dest`'s memory, so the matmul is the whole operation. - a_dest_mat = matricizeview(output_style, a_dest, Val(length(invperm_codomain))) + a_dest_mat = matricizeopview( + output_style, identity, a_dest, invperm_codomain, invperm_domain + ) mul!(a_dest_mat, a1_mat, a2_mat, α, β) elseif iszero(β) # `β` is a strong zero, so `a_dest`'s current data is irrelevant: let the matmul @@ -38,7 +42,9 @@ function contractopadd!( else # `a_dest`'s data contributes through `β`, so gather it, multiply into the gathered # copy, and scatter back. - a_dest_mat = matricizecopy(output_style, a_dest, invperm_codomain, invperm_domain) + a_dest_mat = matricizeopcopy( + output_style, identity, a_dest, invperm_codomain, invperm_domain + ) mul!(a_dest_mat, a1_mat, a2_mat, α, β) unmatricize!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) end diff --git a/src/diagonal.jl b/src/diagonal.jl index 7a465893..0488effa 100644 --- a/src/diagonal.jl +++ b/src/diagonal.jl @@ -44,7 +44,11 @@ end # A `Diagonal` is already a matrix; the `(1 codomain, 1 domain)` matricization is the identity # reshape, so the memory-sharing matricization is `a` itself (keeping it a `Diagonal` for the # `Diagonal`-specialized consumers downstream). -matricizeview(::ReshapeMatricize, a::Diagonal, ::Val{1}) = a +function matricizeopview( + ::ReshapeMatricize, op, a::Diagonal, perm_codomain::Tuple{Int}, perm_domain::Tuple{Int} + ) + return a +end # A `{1,1}` unmatricize (one codomain axis, one domain axis) is the endomorphism identity: the # result stays `Diagonal`, so return `m` directly. The generic `check_input(unmatricize, ...)` # validates the axis lengths against `m`'s size. diff --git a/src/factorizations.jl b/src/factorizations.jl index 5f7903eb..ad1a1ca7 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -16,7 +16,7 @@ using MatrixAlgebraKit: MatrixAlgebraKit # Owned tier: the matrix-level entries mutate their input, so the perm form materializes an # owned matricization following MatrixAlgebraKit's `f(A) = f!(copy_input(f, A))` convention — # a memory-sharing matricization is materialized through `MatrixAlgebraKit.copy_input`, while -# the `matricizecopy` gather is owned by contract and is donated directly (with +# the `matricizeopcopy` gather is owned by contract and is donated directly (with # `copy_input` still applied when the eltype must change) — and the wrapper calls the mutating # entry unconditionally. for f in ( @@ -34,11 +34,16 @@ for f in ( ) ndims(A) == length(perm_codomain) + length(perm_domain) || throw(ArgumentError("Invalid bipermutation")) - A_mat = if ismatricizeview(style, A, perm_codomain, perm_domain) - A_shared = matricizeview(style, A, Val(length(perm_codomain))) + A_mat = + if is_output_view( + matricizeop, style, identity, A, perm_codomain, perm_domain + ) + A_shared = + matricizeopview(style, identity, A, perm_codomain, perm_domain) MatrixAlgebraKit.copy_input(MatrixAlgebraKit.$f, A_shared) else - A_gather = matricizecopy(style, A, perm_codomain, perm_domain) + A_gather = + matricizeopcopy(style, identity, A, perm_codomain, perm_domain) if eltype(A_gather) === float(eltype(A_gather)) A_gather else @@ -56,7 +61,7 @@ for f in ( end # Read-only tier: the matrix-level entries never mutate their input (they copy internally), so -# the perm form consumes the maybe-alias `matricizeperm` matricization directly. +# the perm form consumes the maybe-alias `matricize` matricization directly. for f in ( :gram_eigh_full, :gram_eigh_full_with_pinv, :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, @@ -67,7 +72,7 @@ for f in ( perm_codomain, perm_domain; kwargs... ) - A_mat = matricizeperm(style, A, perm_codomain, perm_domain) + A_mat = matricize(style, A, perm_codomain, perm_domain) F = MatrixAlgebra.$f(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes( map(i -> axes(A, i), (perm_codomain..., perm_domain...)), @@ -174,7 +179,7 @@ function tr(A, ndims_codomain::Val) return tr(MatricizeStyle(A), A, ndims_codomain) end function tr(A, perm_codomain, perm_domain) - return LinearAlgebra.tr(matricizeperm(A, perm_codomain, perm_domain)) + return LinearAlgebra.tr(matricize(A, perm_codomain, perm_domain)) end function tr(A, labels_A, labels_codomain, labels_domain) perm_codomain, perm_domain = @@ -858,14 +863,14 @@ julia> A = randn(2, 3, 2, 3); julia> Id = TensorAlgebra.one(A, (:a, :b, :c, :d), (:a, :b), (:c, :d)); -julia> matricize(Id, Val(2)) ≈ I +julia> matricize(Id, (1, 2), (3, 4)) ≈ I true ``` """ function one end function one!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, ndims_codomain) + A_mat = matricize(style, A, trivialbiperm(A, ndims_codomain)...) MatrixAlgebraKit.one!(A_mat) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, A_mat, axes_codomain, axes_domain) @@ -878,11 +883,14 @@ end # matricization directly when the style declares one at this split, and otherwise fills a # gathered matrix and scatters it back with `unmatricize!`. function one!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - if ismatricizeview(style, A, ndims_codomain) - MatrixAlgebraKit.one!(matricizeview(style, A, ndims_codomain)) + perm_codomain, perm_domain = trivialbiperm(A, ndims_codomain) + if is_output_view(matricizeop, style, identity, A, perm_codomain, perm_domain) + MatrixAlgebraKit.one!( + matricizeopview(style, identity, A, perm_codomain, perm_domain) + ) return A end - A_mat = matricizecopy(style, A, ndims_codomain) + A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) MatrixAlgebraKit.one!(A_mat) return unmatricize!(style, A, A_mat, ndims_codomain) end diff --git a/src/matricize.jl b/src/matricize.jl index 70b2bebb..aa6b3fe4 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -70,116 +70,111 @@ function bipermutedims!( end # ===================================== matricize ======================================== -# Copy convention: `bipermutedims`/`permutedims` always copy (Base `permutedims` semantics). At -# the trivial (`Val`) split the sharing story is exact: `matricizeview` shares `a`'s memory, -# `matricizecopy` returns fresh storage the caller owns, and `matricize` aliases `a` iff -# `ismatricizeview` — so a consumer that writes into a destination checks the trait and writes -# through `matricizeview`, and a consumer that mutates an input either owns a `matricizecopy` -# result by contract or materializes an owned matrix with `MatrixAlgebraKit.copy_input` (see the -# owned tier in `factorizations.jl`). `matricizeperm`/`matricizeopperm` keep maybe-alias -# semantics (the result may view or copy; treat it as read-only) until the planned op/perm-form -# trait lands, and `matricizeview` deliberately has no perm form pending that op/perm-layer -# design. +# A style implements four hooks, all taking the operation, the array and the bipermutation: +# +# `allocate_output(matricizeop, style, op, a, pc, pd)` the matrix destination +# `matricizeop!(dest, style, op, a, pc, pd)` write the matricization into it +# `matricizeopview(style, op, a, pc, pd)` partial: the aliasing form +# `is_output_view(matricizeop, style, op, a, pc, pd)` whether the aliasing form applies +# +# Everything else is derived. `matricizeopcopy` allocates and writes, so it always returns fresh +# storage the caller owns. `matricizeop` returns the view where the style declares one and the copy +# otherwise, i.e. it has maybe-alias semantics and its result must be treated as read-only. +# `matricize` is `matricizeop` at `identity`. +# +# Allocation is a hook rather than generic machinery because computing a matricized destination +# needs the fused axes, which only the style knows: TensorAlgebra deliberately has no generic +# axis-fusion interface. It is also what makes the copy path terminate, since `matricizeop!` is a +# distinct function from the router rather than a re-entry into it. -# `matricize` at the trivial split routes on the style's sharing declaration. Styles implement -# the three leaves (`ismatricizeview`, `matricizeview`, `matricizecopy`) rather than overloading -# `matricize` itself. This assumes the permutation was already performed. -function matricize(style::MatricizeStyle, a, ndims_codomain::Val) - ismatricizeview(style, a, ndims_codomain) && - return matricizeview(style, a, ndims_codomain) - return matricizecopy(style, a, ndims_codomain) +""" + matricizeop(op, a, perm_codomain, perm_domain) + +Matricize `a` across the bipermutation with the element-wise operation `op` folded in, i.e. a +matrix representing `op.(permutedims(a, (perm_codomain..., perm_domain...)))` with the codomain +fused to rows and the domain to columns. + +Has "maybe alias" semantics: the result may share `a`'s memory or be fresh storage, depending on +the style and the array type. Treat it as read-only. Use `matricizeopcopy` for a matrix the caller +owns, and `matricizeopview` (partial) for one guaranteed to alias. +""" +function matricizeop(op, a, perm_codomain, perm_domain) + return matricizeop(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end -function matricize(a, ndims_codomain::Val) - return matricize(MatricizeStyle(a), a, ndims_codomain) +function matricizeop(style::MatricizeStyle, op, a, perm_codomain, perm_domain) + check_biperm(a, perm_codomain, perm_domain) + is_output_view(matricizeop, style, op, a, perm_codomain, perm_domain) && + return matricizeopview(style, op, a, perm_codomain, perm_domain) + return matricizeopcopy(style, op, a, perm_codomain, perm_domain) end -# Partial: defined only where `ismatricizeview` is `true`, and always returns a matricization -# sharing `a`'s memory (the `StridedView` partial-constructor pattern). -function matricizeview(style::MatricizeStyle, a, ndims_codomain::Val) - return throw(MethodError(matricizeview, (style, a, ndims_codomain))) +""" + matricize(a, perm_codomain, perm_domain) + +`matricizeop` at `identity`. Has the same maybe-alias semantics. +""" +function matricize(a, perm_codomain, perm_domain) + return matricizeop(identity, a, perm_codomain, perm_domain) end -# Total: always returns a matricization in fresh storage the caller owns. -function matricizecopy(style::MatricizeStyle, a, ndims_codomain::Val) - return throw(MethodError(matricizecopy, (style, a, ndims_codomain))) +function matricize(style::MatricizeStyle, a, perm_codomain, perm_domain) + return matricizeop(style, identity, a, perm_codomain, perm_domain) end -# `bipermutedims` always copies and `matricize` might return a view, so the result is -# guaranteed to be a copy. -function matricizecopy( - style::MatricizeStyle, a, - perm_codomain, perm_domain - ) - a_perm = bipermutedims(a, perm_codomain, perm_domain) - return matricize(style, a_perm, Val(length(perm_codomain))) +# Total: always fresh storage the caller owns. +function matricizeopcopy(op, a, perm_codomain, perm_domain) + return matricizeopcopy(MatricizeStyle(a), op, a, perm_codomain, perm_domain) +end +function matricizeopcopy(style::MatricizeStyle, op, a, perm_codomain, perm_domain) + check_biperm(a, perm_codomain, perm_domain) + dest = allocate_output(matricizeop, style, op, a, perm_codomain, perm_domain) + return matricizeop!(dest, style, op, a, perm_codomain, perm_domain) end -function matricizeperm( - a, - perm_codomain, perm_domain +# Partial: defined only where `is_output_view` is `true`, and always returns a matricization +# sharing `a`'s memory (the `StridedView` partial-constructor pattern). +function matricizeopview(style::MatricizeStyle, op, a, perm_codomain, perm_domain) + return throw( + MethodError(matricizeopview, (style, op, a, perm_codomain, perm_domain)) ) - return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) end -# Thin wrapper around `matricizeopperm` with identity op — the actual matricization logic -# (and the matricize-style overload point for folding ops into matricization) lives in -# `matricizeopperm`. -function matricizeperm( - style::MatricizeStyle, a, - perm_codomain, perm_domain + +# Required of every style: write the matricization of `a` into `dest`. +function matricizeop!(dest, style::MatricizeStyle, op, a, perm_codomain, perm_domain) + return throw( + MethodError(matricizeop!, (dest, style, op, a, perm_codomain, perm_domain)) ) - return matricizeopperm(style, identity, a, perm_codomain, perm_domain) end -# ================================== matricizeopperm ===================================== - -""" - matricizeopperm(op, a, perm_codomain, perm_domain) - -Matricize `a` with element-wise operation `op` folded in. Returns a matrix representing -`op.(matricizeperm(a, perm_codomain, perm_domain))`. - -Has "maybe alias" semantics: the result may be a view/wrapper aliasing `a` or a fresh -copy, depending on the matricize style and array type. The caller should treat the result -as read-only. -""" -function matricizeopperm( - op, a, perm_codomain, perm_domain +# Required of every style: the matrix destination `matricizeop!` writes into. +function allocate_output( + ::typeof(matricizeop), style::MatricizeStyle, op, a, perm_codomain, perm_domain + ) + return throw( + MethodError( + allocate_output, + (matricizeop, style, op, a, perm_codomain, perm_domain) + ) ) - return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end -# The identity bipermutation is a no-op permute, so `matricize` runs directly on `a` (a view -# for dense, a gather without the extra permute copy for graded); the fast path requires -# `op === identity`, since a plain view cannot carry a fused `op` like `conj`. The result may -# alias `a` and must be treated as read-only, matching the docstring. -function matricizeopperm( - style::MatricizeStyle, op, a, - perm_codomain, perm_domain - ) - check_biperm(a, perm_codomain, perm_domain) - op === identity && isidentitybiperm(perm_codomain, perm_domain) && - return matricize(style, a, Val(length(perm_codomain))) - a_perm_op = permutedimsop(op, a, perm_codomain, perm_domain) - return matricize(style, a_perm_op, Val(length(perm_codomain))) +# The trivial bipermutation for a rank-`N` array split after `ndims_codomain` dimensions. The +# `Val` entry points that remain build it to reach the bipermutation hooks. +function trivialbiperm(a, ndims_codomain::Val{K}) where {K} + return ntuple(identity, Val(K)), ntuple(i -> K + i, Val(ndims(a) - K)) end -# ================================== ismatricizeview ===================================== -# `true` iff `matricize(style, a, ndims_codomain)` shares `a`'s memory (writes to it are writes -# to `a`) — the `isstrided`/`StridedView` pattern (also TensorKit's `has_shared_permute` and -# TensorOperations' `isblasdestination`). Styles overload the `Val` (trivial split) form to -# declare which splits share; the bipermutation form delegates to it at the identity and is -# `false` (fail-safe) everywhere else. A general `ismatricizeview(style, op, a, perm_codomain, -# perm_domain)` form (op and bipermutation view-sets) is planned; these are its -# `op === identity` special cases. -ismatricizeview(::MatricizeStyle, a, ndims_codomain::Val) = false -function ismatricizeview( - style::MatricizeStyle, a, - perm_codomain, perm_domain +# ================================== is_output_view ====================================== +# `true` iff `matricizeop(style, op, a, perm_codomain, perm_domain)` shares `a`'s memory, so that +# writes to the result are writes to `a`. The `isstrided`/`StridedView` pattern, and the same +# question TensorKit asks with `has_shared_permute` and TensorOperations with `isblasdestination`. +# Keyed on the operation like the other function-keyed hooks (`check_input`, `allocate_output`, +# `output_axes`), so the predicate's arguments are exactly the call's arguments. +function is_output_view( + ::typeof(matricizeop), ::MatricizeStyle, op, a, perm_codomain, perm_domain ) - isidentitybiperm(perm_codomain, perm_domain) || return false - return ismatricizeview(style, a, Val(length(perm_codomain))) + return false end -# ==================================== unmatricize ======================================= # Split form: `axes_codomain` and `axes_domain` are the destination axes for the codomain and # domain groups, given codomain-facing (un-dualized), the same convention as `similar_map`. A # matricize style stores the domain axes dualized, so its overload re-dualizes them with `conj` @@ -207,7 +202,7 @@ end # The bipermutation maps the destination's dimension order to the matrix's: `axes(a_dest)` grouped # by it gives the legs in `m`'s order, and the result is permuted back by its inverse. It is not # intrinsically an inverse permutation — the matricized-contraction destination path happens to -# derive it as `invperm(biperm_dest)`, while a `matricizeperm`/`unmatricize!` round trip passes +# derive it as `invperm(biperm_dest)`, while a `matricize`/`unmatricize!` round trip passes # the same forward bipermutation to both. function unmatricize!( a_dest, m, @@ -250,16 +245,32 @@ end # Defaults to ReshapeMatricize, a simple reshape struct ReshapeMatricize <: MatricizeStyle end MatricizeStyle(::Type{<:AbstractArray}) = ReshapeMatricize() -# A dense reshape matricization is a lazy wrapper at any split, so it always shares memory. -ismatricizeview(::ReshapeMatricize, a, ndims_codomain::Val) = true -function matricizeview(::ReshapeMatricize, a, ndims_codomain::Val) - unval(ndims_codomain) ≤ ndims(a) || - throw(ArgumentError("Codomain length exceeds number of dimensions.")) - size_codomain, size_domain = bipartition(size(a), ndims_codomain) +# A dense reshape shares memory only when the data is already in codomain-then-domain order and +# no operation has to be folded in: a reshape can neither reorder nor carry a `conj`. +function is_output_view( + ::typeof(matricizeop), ::ReshapeMatricize, op, a, perm_codomain, perm_domain + ) + return op === identity && isidentitybiperm(perm_codomain, perm_domain) +end +function matricizeopview(::ReshapeMatricize, op, a, perm_codomain, perm_domain) + size_codomain, size_domain = bipartition(size(a), Val(length(perm_codomain))) return reshape(a, (prod(size_codomain), prod(size_domain))) end -function matricizecopy(style::ReshapeMatricize, a, ndims_codomain::Val) - return copy(matricizeview(style, a, ndims_codomain)) +function allocate_output( + ::typeof(matricizeop), ::ReshapeMatricize, op, a, perm_codomain, perm_domain + ) + T = Base.promote_op(op, eltype(a)) + size_codomain = map(i -> size(a, i), perm_codomain) + size_domain = map(i -> size(a, i), perm_domain) + return similar(a, T, (prod(size_codomain), prod(size_domain))) +end +# The destination is a dense matrix, so reshaping it to the permuted tensor shape is a view and +# the permuted-add writes straight through it. +function matricizeop!(dest, ::ReshapeMatricize, op, a, perm_codomain, perm_domain) + perm = (perm_codomain..., perm_domain...) + dest_tensor = reshape(dest, map(i -> size(a, i), perm)) + bipermutedimsopadd!(dest_tensor, op, a, perm_codomain, perm_domain, true, false) + return dest end # The matricized input's rows must be the fused codomain and its columns the fused domain. # `reshape` alone only checks the total element count, so a wrong split with the right total diff --git a/src/matrixfunctions.jl b/src/matrixfunctions.jl index 1747244f..4b5f5244 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -32,7 +32,7 @@ const MATRIX_FUNCTIONS = [ ] # The matrix functions never mutate their input (they allocate their own outputs), so the -# permuted forms consume the maybe-alias `matricizeperm` matricization read-only, skipping +# permuted forms consume the maybe-alias `matricize` matricization read-only, skipping # the eager `bipermutedims` copy at the identity bipermutation. for f in MATRIX_FUNCTIONS @eval begin @@ -53,7 +53,7 @@ for f in MATRIX_FUNCTIONS perm_codomain, perm_domain; kwargs... ) - a_mat = matricizeperm(style, a, perm_codomain, perm_domain) + a_mat = matricize(style, a, perm_codomain, perm_domain) axes_codomain, axes_domain = bipartition_axes( map(i -> axes(a, i), (perm_codomain..., perm_domain...)), Val(length(perm_codomain)) From b9a884c399bd1ee4ba9deff7bbe59a5394e457fc Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 20:54:17 -0400 Subject: [PATCH 10/33] Say identity, not trivial, for the do-nothing bipermutation `identitybiperm` now matches `isidentitybiperm` and sits beside it. The symmetry sense of `trivial` is a different concept, so the two never apply to the same object. Co-Authored-By: Claude Opus 5 (1M context) --- src/bituple.jl | 7 +++++++ src/factorizations.jl | 4 ++-- src/matricize.jl | 6 ------ 3 files changed, 9 insertions(+), 8 deletions(-) diff --git a/src/bituple.jl b/src/bituple.jl index f898ee75..60c78d52 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -83,3 +83,10 @@ function isidentitybiperm(perm_codomain, perm_domain) perm = (perm_codomain..., perm_domain...) return perm == ntuple(identity, length(perm)) end + +# The identity bipermutation for a rank-`N` array split after `ndims_codomain` dimensions, i.e. +# the one `isidentitybiperm` accepts. Transitional: only the `Val` entry points that have yet to be +# removed build it, to reach the bipermutation hooks. +function identitybiperm(a, ndims_codomain::Val{K}) where {K} + return ntuple(identity, Val(K)), ntuple(i -> K + i, Val(ndims(a) - K)) +end diff --git a/src/factorizations.jl b/src/factorizations.jl index ad1a1ca7..67dd1e8f 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -870,7 +870,7 @@ true function one end function one!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, trivialbiperm(A, ndims_codomain)...) + A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) MatrixAlgebraKit.one!(A_mat) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, A_mat, axes_codomain, axes_domain) @@ -883,7 +883,7 @@ end # matricization directly when the style declares one at this split, and otherwise fills a # gathered matrix and scatters it back with `unmatricize!`. function one!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - perm_codomain, perm_domain = trivialbiperm(A, ndims_codomain) + perm_codomain, perm_domain = identitybiperm(A, ndims_codomain) if is_output_view(matricizeop, style, identity, A, perm_codomain, perm_domain) MatrixAlgebraKit.one!( matricizeopview(style, identity, A, perm_codomain, perm_domain) diff --git a/src/matricize.jl b/src/matricize.jl index aa6b3fe4..3bc4acfb 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -157,12 +157,6 @@ function allocate_output( ) end -# The trivial bipermutation for a rank-`N` array split after `ndims_codomain` dimensions. The -# `Val` entry points that remain build it to reach the bipermutation hooks. -function trivialbiperm(a, ndims_codomain::Val{K}) where {K} - return ntuple(identity, Val(K)), ntuple(i -> K + i, Val(ndims(a) - K)) -end - # ================================== is_output_view ====================================== # `true` iff `matricizeop(style, op, a, perm_codomain, perm_domain)` shares `a`'s memory, so that # writes to the result are writes to `a`. The `isstrided`/`StridedView` pattern, and the same From facedf4f85eb69ba9e7f15c1d4592ac131dd0305 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 21:33:07 -0400 Subject: [PATCH 11/33] Move the tests onto the matricize hooks Also migrates five Val-forwarding calls in the factorization wrappers that a literal-only search had missed, and the second custom style in the factorization tests. Co-Authored-By: Claude Opus 5 (1M context) --- ext/TensorAlgebraTensorKitExt.jl | 57 ++++++++++----------- src/factorizations.jl | 10 ++-- test/test_basics.jl | 50 +++++++++---------- test/test_diagonal.jl | 2 +- test/test_exports.jl | 5 +- test/test_factorizations.jl | 69 +++++++++++++++++++------- test/test_matricize.jl | 85 ++++++++++++++++++-------------- test/test_matricizestyle.jl | 28 ++++++++--- test/test_tensorkitext.jl | 2 +- 9 files changed, 181 insertions(+), 127 deletions(-) diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index aa770db2..58fa7093 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -243,44 +243,41 @@ end struct TensorKitMatricize <: TensorAlgebra.MatricizeStyle end TensorAlgebra.MatricizeStyle(::Type{<:AbstractTensorMap}) = TensorKitMatricize() -# `permute` at the tensor's own codomain/domain split is trivial and returns `t` itself, so -# the matching split is the one memory-sharing matricization (TensorKit's own -# `has_shared_permute` notion); any other split regroups into a fresh `TensorMap`. -function TensorAlgebra.ismatricizeview( - ::TensorKitMatricize, ::AbstractTensorMap{<:Any, <:Any, K}, ::Val{K} - ) where {K} - return true +# `permute` at the tensor's own codomain/domain split is trivial and returns `t` itself, so that +# split is the one memory-sharing matricization (TensorKit's own `has_shared_permute` notion). Any +# other split, or a folded `conj`, regroups into a fresh `TensorMap`. +function TensorAlgebra.is_output_view( + ::typeof(TensorAlgebra.matricizeop), ::TensorKitMatricize, op, + t::AbstractTensorMap, perm_codomain, perm_domain + ) + return op === identity && + TensorAlgebra.isidentitybiperm(perm_codomain, perm_domain) && + length(perm_codomain) == numout(t) end -TensorAlgebra.ismatricizeview(::TensorKitMatricize, ::AbstractTensorMap, ::Val) = false -function TensorAlgebra.matricizeview( - ::TensorKitMatricize, t::AbstractTensorMap{<:Any, <:Any, K}, ::Val{K} - ) where {K} +function TensorAlgebra.matricizeopview( + ::TensorKitMatricize, op, t::AbstractTensorMap, perm_codomain, perm_domain + ) return t end -function TensorAlgebra.matricizecopy( - ::TensorKitMatricize, t::AbstractTensorMap, ndims_codomain::Val{K} - ) where {K} - N = numind(t) - return permute( - t, - (ntuple(identity, Val(K)), ntuple(i -> K + i, Val(N - K))); - copy = true +# A `TensorMap`'s matricization is a regrouping of its indices, so the destination is a `TensorMap` +# over the regrouped space and the write is the ordinary permuted-add. `bipermutedimsopadd!` above +# routes that through `tensoradd!`, which realizes the permutation, the `op === conj` conjugation +# and the scaling in one call, so no separate handling of `op` is needed here. +function TensorAlgebra.allocate_output( + ::typeof(TensorAlgebra.matricizeop), ::TensorKitMatricize, op, + t::AbstractTensorMap, perm_codomain, perm_domain ) + return similar(t, permute(space(t), (perm_codomain, perm_domain))) end - -# The identity fill on the regrouped map is TensorKit's own `one!` (MatrixAlgebraKit's -# `one!` speaks `AbstractMatrix` only). -function TensorAlgebra.one!!( - style::TensorKitMatricize, A::AbstractTensorMap, ndims_codomain::Val; kwargs... +function TensorAlgebra.matricizeop!( + dest::AbstractTensorMap, ::TensorKitMatricize, op, + t::AbstractTensorMap, perm_codomain, perm_domain + ) + return TensorAlgebra.bipermutedimsopadd!( + dest, op, t, perm_codomain, perm_domain, true, false ) - return TensorKit.one!(TensorAlgebra.matricize(style, A, ndims_codomain)) end -# `unmatricize` reconstructs the codomain/domain axes from the matrix `m`. A `TensorMap` already -# is the linear map its space describes, so the only valid request is the one whose codomain/domain -# split matches `m`'s own space, and `unmatricize` returns `m` unchanged. The domain axes arrive -# codomain-facing (un-dualized), which is exactly TensorKit's domain convention, so they build the -# domain `ProductSpace` directly. function TensorAlgebra.unmatricize( ::TensorKitMatricize, m::AbstractTensorMap, axes_codomain, axes_domain ) diff --git a/src/factorizations.jl b/src/factorizations.jl index 67dd1e8f..83bb9980 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -173,7 +173,7 @@ true ``` """ function tr(style::MatricizeStyle, A, ndims_codomain::Val) - return LinearAlgebra.tr(matricize(style, A, ndims_codomain)) + return LinearAlgebra.tr(matricize(style, A, identitybiperm(A, ndims_codomain)...)) end function tr(A, ndims_codomain::Val) return tr(MatricizeStyle(A), A, ndims_codomain) @@ -559,7 +559,7 @@ The output satisfies `N' * A ≈ 0` and `N' * N ≈ I`. left_null function left_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, ndims_codomain) + A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) N = MatrixAlgebraKit.left_null!(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return unmatricize(style, N, axes_codomain, (conj(axes(N, ndims(N))),)) @@ -596,7 +596,7 @@ The output satisfies `A * Nᴴ' ≈ 0` and `Nᴴ * Nᴴ' ≈ I`. right_null function right_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, ndims_codomain) + A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) Nᴴ = MatrixAlgebraKit.right_null!(A_mat; kwargs...) _, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, Nᴴ, (axes(Nᴴ, 1),), axes_domain) @@ -652,7 +652,7 @@ gram_eigh_full function gram_eigh_full!!( style::MatricizeStyle, A, ndims_codomain::Val; kwargs... ) - A_mat = matricize(style, A, ndims_codomain) + A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) X = MatrixAlgebra.gram_eigh_full!!(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)) @@ -711,7 +711,7 @@ gram_eigh_full_with_pinv function gram_eigh_full_with_pinv!!( style::MatricizeStyle, A, ndims_codomain::Val; kwargs... ) - A_mat = matricize(style, A, ndims_codomain) + A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) X, Y = MatrixAlgebra.gram_eigh_full_with_pinv!!(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)), diff --git a/test/test_basics.jl b/test/test_basics.jl index 8f45811d..b0a119cb 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,7 +1,7 @@ import TensorAlgebra using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, - contract!, contractadd!, length_codomain, length_domain, matricizeperm, unmatricize, + contract!, contractadd!, length_codomain, length_domain, matricize, unmatricize, unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -53,68 +53,68 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @testset "matricize (eltype=$elt)" for elt in elts a = randn(elt, 2, 3, 4, 5) - a_fused = matricizeperm(a, (1, 2), (3, 4)) + a_fused = matricize(a, (1, 2), (3, 4)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, 6, 20) - a_fused = matricizeperm(a, (3, 1), (2, 4)) + a_fused = matricize(a, (3, 1), (2, 4)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(permutedims(a, (3, 1, 2, 4)), (8, 15)) - a_fused = matricizeperm(a, (3, 1, 2), (4,)) + a_fused = matricize(a, (3, 1, 2), (4,)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(permutedims(a, (3, 1, 2, 4)), (24, 5)) - a_fused = matricizeperm(a, (2, 4), (3, 1)) + a_fused = matricize(a, (2, 4), (3, 1)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(permutedims(a, (2, 4, 3, 1)), (15, 8)) # Degenerate splits: everything in the domain, then everything in the codomain. - a_fused = matricizeperm(a, (), (1, 2, 3, 4)) + a_fused = matricize(a, (), (1, 2, 3, 4)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, (1, 120)) - a_fused = matricizeperm(a, (1, 2, 3, 4), ()) + a_fused = matricize(a, (1, 2, 3, 4), ()) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, (120, 1)) - @test_throws MethodError matricizeperm(a, (1, 2), (3,), (4,)) - @test_throws MethodError matricizeperm(a, (1, 2, 3, 4)) - @test_throws ArgumentError matricizeperm(a, (1, 2), (3,)) + @test_throws MethodError matricize(a, (1, 2), (3,), (4,)) + @test_throws MethodError matricize(a, (1, 2, 3, 4)) + @test_throws ArgumentError matricize(a, (1, 2), (3,)) v = ones(elt, 2) - a_fused = matricizeperm(v, (1,), ()) + a_fused = matricize(v, (1,), ()) @test eltype(a_fused) === elt @test a_fused ≈ ones(elt, 2, 1) - a_fused = matricizeperm(v, (), (1,)) + a_fused = matricize(v, (), (1,)) @test eltype(a_fused) === elt @test a_fused ≈ ones(elt, 1, 2) - a_fused = matricizeperm(ones(elt), (), ()) + a_fused = matricize(ones(elt), (), ()) @test eltype(a_fused) === elt @test a_fused ≈ ones(elt, 1, 1) end - @testset "matricizeopperm (eltype=$elt)" for elt in elts + @testset "matricizeop (eltype=$elt)" for elt in elts rng = StableRNG(123) a = randn(rng, elt, 2, 3, 4) # identity op: should match matricize exactly - m = TensorAlgebra.matricizeopperm(identity, a, (1,), (2, 3)) - m_ref = matricizeperm(a, (1,), (2, 3)) + m = TensorAlgebra.matricizeop(identity, a, (1,), (2, 3)) + m_ref = matricize(a, (1,), (2, 3)) @test m ≈ m_ref - m = TensorAlgebra.matricizeopperm(identity, a, (3, 1), (2,)) - m_ref = matricizeperm(a, (3, 1), (2,)) + m = TensorAlgebra.matricizeop(identity, a, (3, 1), (2,)) + m_ref = matricize(a, (3, 1), (2,)) @test m ≈ m_ref - m = TensorAlgebra.matricizeopperm(identity, a, (2, 3), (1,)) - m_ref = matricizeperm(a, (2, 3), (1,)) + m = TensorAlgebra.matricizeop(identity, a, (2, 3), (1,)) + m_ref = matricize(a, (2, 3), (1,)) @test m ≈ m_ref # conj op - m = TensorAlgebra.matricizeopperm(conj, a, (1,), (2, 3)) - m_ref = conj.(matricizeperm(a, (1,), (2, 3))) + m = TensorAlgebra.matricizeop(conj, a, (1,), (2, 3)) + m_ref = conj.(matricize(a, (1,), (2, 3))) @test m ≈ m_ref - m = TensorAlgebra.matricizeopperm(conj, a, (3, 1), (2,)) - m_ref = conj.(matricizeperm(a, (3, 1), (2,))) + m = TensorAlgebra.matricizeop(conj, a, (3, 1), (2,)) + m_ref = conj.(matricize(a, (3, 1), (2,))) @test m ≈ m_ref end @@ -136,7 +136,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int unmatricize!(a, m, (1, 2), (3, 4)) @test a ≈ a0 - m1 = matricizeperm(a0, perm_codomain, perm_domain) + m1 = matricize(a0, perm_codomain, perm_domain) a = similar(a0) unmatricize!(a, m1, perm_codomain, perm_domain) @test a ≈ a0 diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index c492918a..45fcefe8 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -41,7 +41,7 @@ using Test: @test, @test_throws, @testset end @testset "matricize(1, 1) is the identity reshape" begin - m = TensorAlgebra.matricize(TensorAlgebra.ReshapeMatricize(), d, Val(1)) + m = TensorAlgebra.matricize(TensorAlgebra.ReshapeMatricize(), d, (1,), (2,)) @test m === d end diff --git a/test/test_exports.jl b/test/test_exports.jl index 5bab6ad0..0ff6b33e 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -43,8 +43,9 @@ using Test: @test, @testset :biperm, :bipartition, :cat_similar, :concatenate, :concatenate!, :ContractAlgorithm, :contractopadd!, :data, :datatype, :directsum, - :flattenlinear, :label_type, - :matricizeopperm, :permutedims, :permutedims!, :scalar, :similar_map, + :flattenlinear, :is_output_view, :label_type, + :matricize, :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, + :permutedims, :permutedims!, :scalar, :similar_map, :TensorOperationsAlgorithm, :to_range, :tr, :tryflattenlinear, :ungrade, :zero!, :scale!, :permuteddims, :PermutedDims, diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index f8934459..d69a007d 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -7,6 +7,9 @@ using TensorAlgebra: TensorAlgebra, contract, eig_full, eig_vals, eigh_full, eig using Test: @test, @testset using TestExtras: @constinferred +# Matricize without permuting: the identity bipermutation for a split after `k` dimensions. +splitperms(a, k) = (ntuple(identity, k), ntuple(i -> k + i, ndims(a) - k)) + elts = (Float64, ComplexF64) # QR Decomposition @@ -373,7 +376,7 @@ end @test size(Id) == size(A) @test eltype(Id) === T - @test TensorAlgebra.matricize(Id, Val(2)) ≈ I + @test TensorAlgebra.matricize(Id, splitperms(Id, 2)...) ≈ I # `Val`, perm, and label entries agree. @test TensorAlgebra.one(A, Val(2)) ≈ Id @@ -386,7 +389,7 @@ end B = randn(T, 2, 2, 2, 2) labels_B = (:a, :c, :b, :d) Id_perm = TensorAlgebra.one(B, labels_B, labels_cod, labels_dom) - @test TensorAlgebra.matricize(Id_perm, Val(2)) ≈ I + @test TensorAlgebra.matricize(Id_perm, splitperms(Id_perm, 2)...) ≈ I # Perm- and biperm-tuple forms agree with the label form. @test TensorAlgebra.one(B, (1, 3), (2, 4)) ≈ Id_perm @@ -394,12 +397,12 @@ end C = randn(T, 2, 3, 2, 3) Cret = @constinferred TensorAlgebra.one!(C, Val(2)) @test Cret === C - @test TensorAlgebra.matricize(C, Val(2)) ≈ I + @test TensorAlgebra.matricize(C, splitperms(C, 2)...) ≈ I @test C ≈ TensorAlgebra.one(A, Val(2)) # `unmatricize!` scatters a fused matrix back into an existing array. D = randn(T, 2, 3, 2, 3) - Dmat = TensorAlgebra.matricize(D, Val(2)) + Dmat = TensorAlgebra.matricize(D, splitperms(D, 2)...) E = similar(D) Eret = TensorAlgebra.unmatricize!(E, Dmat, Val(2)) @test Eret === E @@ -430,23 +433,25 @@ end (((1, 2), (3,)), ((3, 1), (2,)), ((3,), (1, 2)), ((2,), (3, 1))) k = length(perm_codomain) A_perm = TensorAlgebra.bipermutedims(A, perm_codomain, perm_domain) - A_mat = TensorAlgebra.matricize(A_perm, Val(k)) + A_mat = TensorAlgebra.matricize(A_perm, splitperms(A_perm, k)...) for f in (qr_compact, lq_compact, left_orth, right_orth) X, Y = f(A, perm_codomain, perm_domain) - @test TensorAlgebra.matricize(X, Val(k)) * - TensorAlgebra.matricize(Y, Val(1)) ≈ A_mat + @test TensorAlgebra.matricize(X, splitperms(X, k)...) * + TensorAlgebra.matricize(Y, splitperms(Y, 1)...) ≈ A_mat end for f in (svd_compact, svd_trunc) U, S, Vᴴ = f(A, perm_codomain, perm_domain) - U_mat = TensorAlgebra.matricize(U, Val(k)) - @test U_mat * S * TensorAlgebra.matricize(Vᴴ, Val(1)) ≈ A_mat + U_mat = TensorAlgebra.matricize(U, splitperms(U, k)...) + @test U_mat * S * TensorAlgebra.matricize(Vᴴ, splitperms(Vᴴ, 1)...) ≈ A_mat @test U_mat' * U_mat ≈ I end @test svd_vals(A, perm_codomain, perm_domain) ≈ LinearAlgebra.svdvals(A_mat) - N = TensorAlgebra.matricize(left_null(A, perm_codomain, perm_domain), Val(k)) + N_tensor = left_null(A, perm_codomain, perm_domain) + N = TensorAlgebra.matricize(N_tensor, splitperms(N_tensor, k)...) @test norm(N' * A_mat) ≈ 0 atol = 1.0e-13 @test N' * N ≈ I - Nᴴ = TensorAlgebra.matricize(right_null(A, perm_codomain, perm_domain), Val(1)) + Nᴴ_tensor = right_null(A, perm_codomain, perm_domain) + Nᴴ = TensorAlgebra.matricize(Nᴴ_tensor, splitperms(Nᴴ_tensor, 1)...) @test norm(A_mat * Nᴴ') ≈ 0 atol = 1.0e-13 @test Nᴴ * Nᴴ' ≈ I @test A == Acopy @@ -456,9 +461,9 @@ end for (perm_codomain, perm_domain) in (((1, 2), (3, 4)), ((3, 4), (1, 2)), ((2, 3), (4, 1))) B_perm = TensorAlgebra.bipermutedims(B, perm_codomain, perm_domain) - B_mat = Matrix(TensorAlgebra.matricize(B_perm, Val(2))) + B_mat = Matrix(TensorAlgebra.matricize(B_perm, splitperms(B_perm, 2)...)) D, V = eig_full(B, perm_codomain, perm_domain) - V_mat = TensorAlgebra.matricize(V, Val(2)) + V_mat = TensorAlgebra.matricize(V, splitperms(V, 2)...) @test B_mat * V_mat ≈ V_mat * D sortvals(v) = sort(v; by = x -> (real(x), imag(x))) @test sortvals(eig_vals(B, perm_codomain, perm_domain)) ≈ @@ -481,11 +486,37 @@ module FactorizationMatricizeTestUtils end struct AliasingMatricize <: TA.MatricizeStyle end TA.MatricizeStyle(::Type{<:AliasingArray}) = AliasingMatricize() - function TA.matricize(::AliasingMatricize, a::AliasingArray, ndims_codomain::Val) - return TA.matricize(TA.ReshapeMatricize(), a.parent, ndims_codomain) + # Delegate every hook to the dense style on the unwrapped parent, so the matricization + # aliases exactly where a plain `Array`'s would. + unwrap(a::AliasingArray) = a.parent + unwrap(a::AbstractArray) = a + function TA.is_output_view( + ::typeof(TA.matricizeop), ::AliasingMatricize, op, a, perm_codomain, perm_domain + ) + return TA.is_output_view( + TA.matricizeop, TA.ReshapeMatricize(), op, unwrap(a), perm_codomain, perm_domain + ) + end + function TA.matricizeopview( + ::AliasingMatricize, op, a, perm_codomain, perm_domain + ) + return TA.matricizeopview( + TA.ReshapeMatricize(), op, unwrap(a), perm_codomain, perm_domain + ) end - function TA.matricize(::AliasingMatricize, a::AbstractArray, ndims_codomain::Val) - return TA.matricize(TA.ReshapeMatricize(), a, ndims_codomain) + function TA.allocate_output( + ::typeof(TA.matricizeop), ::AliasingMatricize, op, a, perm_codomain, perm_domain + ) + return TA.allocate_output( + TA.matricizeop, TA.ReshapeMatricize(), op, unwrap(a), perm_codomain, perm_domain + ) + end + function TA.matricizeop!( + dest, ::AliasingMatricize, op, a, perm_codomain, perm_domain + ) + return TA.matricizeop!( + dest, TA.ReshapeMatricize(), op, unwrap(a), perm_codomain, perm_domain + ) end function TA.unmatricize(::AliasingMatricize, m, axes_codomain, axes_domain) return AliasingArray( @@ -512,9 +543,9 @@ end parent_copy = copy(parent) Q, R = qr_compact(A, Val(1)) @test parent == parent_copy - Q_mat = TensorAlgebra.matricize(Q, Val(1)) + Q_mat = TensorAlgebra.matricize(Q, splitperms(Q, 1)...) @test eltype(Q_mat) === Float64 - @test Q_mat * TensorAlgebra.matricize(R, Val(1)) ≈ reshape(parent, 2, 12) + @test Q_mat * TensorAlgebra.matricize(R, splitperms(R, 1)...) ≈ reshape(parent, 2, 12) @test svd_vals(A, (1,), (2, 3)) ≈ LinearAlgebra.svdvals(reshape(float.(parent), 2, 12)) @test parent == parent_copy end diff --git a/test/test_matricize.jl b/test/test_matricize.jl index d16eba71..0582b4bd 100644 --- a/test/test_matricize.jl +++ b/test/test_matricize.jl @@ -1,6 +1,6 @@ using StableRNGs: StableRNG -using TensorAlgebra: TensorAlgebra, ReshapeMatricize, ismatricizeview, matricize, - matricizecopy, matricizeopperm, matricizeperm, matricizeview +using TensorAlgebra: TensorAlgebra, ReshapeMatricize, is_output_view, matricize, + matricizeop, matricizeop!, matricizeopcopy, matricizeopview using Test: @test, @test_throws, @testset # A non-`ReshapeMatricize` style, to check the always-safe generic fallback. @@ -14,86 +14,95 @@ function matricize_ref(a, perm_codomain, perm_domain) return reshape(a_perm, (nrow, ncol)) end -@testset "maybe-view matricizeperm (eltype=$elt)" for elt in (Float64, ComplexF64) +@testset "maybe-view matricize (eltype=$elt)" for elt in (Float64, ComplexF64) a = randn(StableRNG(123), elt, 2, 3, 4) # Identity bipermutation: correct values and a view aliasing `a`. - m = matricizeperm(a, (1,), (2, 3)) + m = matricize(a, (1,), (2, 3)) @test m ≈ matricize_ref(a, (1,), (2, 3)) @test Base.mightalias(m, a) # Every other bipermutation is a fresh permuted copy in matricized layout (no lazy # wrappers), including the codomain/domain swap. for (pc, pd) in (((2, 3), (1,)), ((3, 1), (2,))) - m = matricizeperm(a, pc, pd) + m = matricize(a, pc, pd) @test m ≈ matricize_ref(a, pc, pd) @test m isa Matrix @test !Base.mightalias(m, a) end - @test_throws ArgumentError matricizeperm(a, (1,), (2,)) + @test_throws ArgumentError matricize(a, (1,), (2,)) # `conj` cannot ride a view, so it copies even on the identity bipermutation. - m = matricizeopperm(conj, a, (1,), (2, 3)) + m = matricizeop(conj, a, (1,), (2, 3)) @test m ≈ conj.(matricize_ref(a, (1,), (2, 3))) @test !Base.mightalias(m, a) - m = matricizeopperm(conj, a, (2, 3), (1,)) + m = matricizeop(conj, a, (2, 3), (1,)) @test m ≈ conj.(matricize_ref(a, (2, 3), (1,))) @test !Base.mightalias(m, a) end -@testset "ismatricizeview" begin +@testset "is_output_view" begin a = randn(StableRNG(321), 2, 3, 4) style = ReshapeMatricize() - # A dense reshape matricization shares memory at every trivial split. - @test ismatricizeview(style, a, Val(1)) - @test ismatricizeview(style, a, (1,), (2, 3)) + # A dense reshape shares memory at the identity bipermutation. + @test is_output_view(matricizeop, style, identity, a, (1,), (2, 3)) + @test is_output_view(matricizeop, style, identity, a, (), (1, 2, 3)) + @test is_output_view(matricizeop, style, identity, a, (1, 2, 3), ()) - # The bipermutation form declares sharing only at the identity: a swap or interleaving - # bipermutation routes through the consumers' gather branches. - @test !ismatricizeview(style, a, (2, 3), (1,)) - @test !ismatricizeview(style, a, (3, 1), (2,)) + # Not at a swap or an interleaving, which route through the gather branch instead. + @test !is_output_view(matricizeop, style, identity, a, (2, 3), (1,)) + @test !is_output_view(matricizeop, style, identity, a, (3, 1), (2,)) + + # And never with an operation folded in, since a reshape cannot carry a `conj`. + @test !is_output_view(matricizeop, style, conj, a, (1,), (2, 3)) # A generic style declares nothing (fail-safe default). - @test !ismatricizeview(DummyMatricize(), a, Val(1)) - @test !ismatricizeview(DummyMatricize(), a, (1,), (2, 3)) + @test !is_output_view(matricizeop, DummyMatricize(), identity, a, (1,), (2, 3)) # Writes to the shared matricization are writes to `a`. - m = matricizeview(style, a, Val(1)) + m = matricizeopview(style, identity, a, (1,), (2, 3)) @test m == matricize_ref(a, (1,), (2, 3)) m[1, 1] = 42 @test a[1, 1, 1] == 42 end -@testset "ismatricizeview coherence" begin +@testset "is_output_view coherence" begin rng = StableRNG(11) a = randn(rng, 2, 3, 4) style = ReshapeMatricize() - # A declared share means `matricizeview` (and so `matricize`) aliases `a`, while - # `matricizecopy` never does. + # A declared share means `matricizeopview` (and so `matricize`) aliases `a`, while + # `matricizeopcopy` never does. for K in 0:3 - if ismatricizeview(style, a, Val(K)) - m = matricizeview(style, a, Val(K)) - @test Base.mightalias(m, a) - @test matricize(style, a, Val(K)) == m - end - @test !Base.mightalias(matricizecopy(style, a, Val(K)), a) - - # The perm form of the copy leaf is owned too, and at the trivial bipermutation it - # matches the `Val` form. pc = ntuple(identity, K) pd = ntuple(i -> K + i, 3 - K) - m_perm = matricizecopy(style, a, pc, pd) - @test !Base.mightalias(m_perm, a) - @test m_perm == matricizecopy(style, a, Val(K)) + if is_output_view(matricizeop, style, identity, a, pc, pd) + m = matricizeopview(style, identity, a, pc, pd) + @test Base.mightalias(m, a) + @test matricize(style, a, pc, pd) == m + end + m_copy = matricizeopcopy(style, identity, a, pc, pd) + @test !Base.mightalias(m_copy, a) + @test m_copy == matricize_ref(a, pc, pd) end for (pc, pd) in (((2, 3), (1,)), ((3, 1), (2,))) - m = matricizecopy(style, a, pc, pd) + m = matricizeopcopy(style, identity, a, pc, pd) @test m ≈ matricize_ref(a, pc, pd) @test !Base.mightalias(m, a) end - @test_throws ArgumentError matricizecopy(style, a, (1,), (2,)) + @test_throws ArgumentError matricizeopcopy(style, identity, a, (1,), (2,)) + + # The allocation and write hooks compose into the copy form. + for (pc, pd) in (((1,), (2, 3)), ((3, 1), (2,))) + for op in (identity, conj) + dest = TensorAlgebra.allocate_output(matricizeop, style, op, a, pc, pd) + @test size(dest) == size(matricize_ref(a, pc, pd)) + matricizeop!(dest, style, op, a, pc, pd) + @test dest ≈ op.(matricize_ref(a, pc, pd)) + @test dest ≈ matricizeopcopy(style, op, a, pc, pd) + end + end # Both destination branches of a consumer (`contractadd!`) behave: the shared-view route # for the identity destination bipermutation and the gather/scatter route otherwise. @@ -123,14 +132,14 @@ end # Identity-bipermutation view tracks an in-place update of `a`. a = randn(rng, 2, 3, 4) - m = matricizeperm(a, (1,), (2, 3)) + m = matricize(a, (1,), (2, 3)) a .= randn(rng, 2, 3, 4) @test m ≈ matricize_ref(a, (1,), (2, 3)) # Permuted copies are independent of later updates to `a`. for (pc, pd) in (((2, 3), (1,)), ((3, 1), (2,))) a = randn(rng, 2, 3, 4) - m = matricizeperm(a, pc, pd) + m = matricize(a, pc, pd) snapshot = copy(m) a .= a .+ 1 @test m == snapshot diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index 8d13c40b..3c623e2d 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -9,12 +9,28 @@ module MatricizeStyleTestUtils end struct MyArrayMatricize <: TA.MatricizeStyle end TA.MatricizeStyle(::Type{<:MyArray}) = MyArrayMatricize() - # Minimal fold/unfold leaves so a round-trip (`one!`) can run through the custom style: - # both dispatch on `MyArrayMatricize`, so an unfold whose style was re-derived from the - # plain fused matrix instead of threaded through would miss them and error. - TA.ismatricizeview(::MyArrayMatricize, a, ::Val) = false - function TA.matricizecopy(::MyArrayMatricize, a::MyArray, ndims_codomain::Val) - return TA.matricizecopy(TA.ReshapeMatricize(), a.parent, ndims_codomain) + # Minimal hooks so a round trip (`one!`) can run through the custom style. All of them + # dispatch on `MyArrayMatricize`, so a path whose style was re-derived from the plain fused + # matrix instead of threaded through would miss them and error. + function TA.is_output_view( + ::typeof(TA.matricizeop), ::MyArrayMatricize, op, a, perm_codomain, perm_domain + ) + return false + end + function TA.allocate_output( + ::typeof(TA.matricizeop), ::MyArrayMatricize, op, a::MyArray, + perm_codomain, perm_domain + ) + return TA.allocate_output( + TA.matricizeop, TA.ReshapeMatricize(), op, a.parent, perm_codomain, perm_domain + ) + end + function TA.matricizeop!( + dest, ::MyArrayMatricize, op, a::MyArray, perm_codomain, perm_domain + ) + return TA.matricizeop!( + dest, TA.ReshapeMatricize(), op, a.parent, perm_codomain, perm_domain + ) end function TA.unmatricize!( ::MyArrayMatricize, a_dest::MyArray, m, diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index 557b4983..30b36a27 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -60,7 +60,7 @@ using Test: @test, @test_throws, @testset # `unmatricize` takes the domain axes codomain-facing (un-dualized), so pass `B`, `C1` # directly rather than the dualized `space(t, 3)`, `space(t, 4)`. axes_domain = (B, C1) - m = matricize(t, Val(2)) + m = matricize(t, (1, 2), (3, 4)) @test space(m) == space(t) back = unmatricize(m, axes_codomain, axes_domain) @test back ≈ t From 21eb4618f4b6a5098a654be09e969920fcedecc7 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 22:35:37 -0400 Subject: [PATCH 12/33] Offer the split-only spelling as convenience again Matricizing without permuting is worth a short spelling. Removing `Val` as a dispatch tier was what mattered: a style implements the bipermutation hooks and never these, so the copy path cannot recurse through the router. Co-Authored-By: Claude Opus 5 (1M context) --- src/bituple.jl | 4 ++-- src/matricize.jl | 16 ++++++++++++++++ 2 files changed, 18 insertions(+), 2 deletions(-) diff --git a/src/bituple.jl b/src/bituple.jl index 60c78d52..d9344a55 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -85,8 +85,8 @@ function isidentitybiperm(perm_codomain, perm_domain) end # The identity bipermutation for a rank-`N` array split after `ndims_codomain` dimensions, i.e. -# the one `isidentitybiperm` accepts. Transitional: only the `Val` entry points that have yet to be -# removed build it, to reach the bipermutation hooks. +# the one `isidentitybiperm` accepts. The split-only `Val` conveniences build it to reach the +# bipermutation forms. function identitybiperm(a, ndims_codomain::Val{K}) where {K} return ntuple(identity, Val(K)), ntuple(i -> K + i, Val(ndims(a) - K)) end diff --git a/src/matricize.jl b/src/matricize.jl index 3bc4acfb..23cff0f2 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -120,6 +120,22 @@ function matricize(style::MatricizeStyle, a, perm_codomain, perm_domain) return matricizeop(style, identity, a, perm_codomain, perm_domain) end +# Split-only convenience: matricize after `ndims_codomain` dimensions without permuting. Sugar over +# the bipermutation forms, not a dispatch tier. A style implements the hooks above and never these, +# which is what keeps the copy path from recursing back through the router. +function matricize(a, ndims_codomain::Val) + return matricize(a, identitybiperm(a, ndims_codomain)...) +end +function matricize(style::MatricizeStyle, a, ndims_codomain::Val) + return matricize(style, a, identitybiperm(a, ndims_codomain)...) +end +function matricizeop(op, a, ndims_codomain::Val) + return matricizeop(op, a, identitybiperm(a, ndims_codomain)...) +end +function matricizeop(style::MatricizeStyle, op, a, ndims_codomain::Val) + return matricizeop(style, op, a, identitybiperm(a, ndims_codomain)...) +end + # Total: always fresh storage the caller owns. function matricizeopcopy(op, a, perm_codomain, perm_domain) return matricizeopcopy(MatricizeStyle(a), op, a, perm_codomain, perm_domain) From 9e41545562efe649166e93858ac13e244dce57d7 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 22:48:42 -0400 Subject: [PATCH 13/33] Remove gram_eigh_full, gram_eigh_full_with_pinv, and sqrth_invsqrth_safe The gram factorizations were only used by higher-level network code, so they belong in the package that needs them. `sqrth_invsqrth_safe` only saves an eigendecomposition over calling the two separately. Co-Authored-By: Claude Opus 5 (1M context) --- src/MatrixAlgebra.jl | 121 +-------------------------- src/TensorAlgebra.jl | 4 +- src/factorizations.jl | 158 ++---------------------------------- test/test_exports.jl | 6 -- test/test_factorizations.jl | 58 +------------ test/test_matrixalgebra.jl | 36 -------- 6 files changed, 11 insertions(+), 372 deletions(-) diff --git a/src/MatrixAlgebra.jl b/src/MatrixAlgebra.jl index 82aeece6..05e5aba6 100644 --- a/src/MatrixAlgebra.jl +++ b/src/MatrixAlgebra.jl @@ -1,15 +1,12 @@ module MatrixAlgebra -export gram_eigh_full, - gram_eigh_full_with_pinv, - invsqrt_diag_safe, +export invsqrt_diag_safe, invsqrth_safe, pow_diag_safe, pow_diag_safe!, powh_safe, sqrt_diag_safe, - sqrth_safe, - sqrth_invsqrth_safe + sqrth_safe using LinearAlgebra: LinearAlgebra, Diagonal, isdiag, norm using MatrixAlgebraKit: MatrixAlgebraKit as MAK @@ -165,120 +162,6 @@ $(_clamp_kwargs_doc("M")) """ invsqrth_safe(M; kwargs...) = powh_safe(M, -1 // 2; kwargs...) -""" - sqrth_invsqrth_safe(M; alg=nothing, atol=0, rtol=eps(real(eltype(M)))^(3//4)) -> M^(1//2), M^(-1//2) - -Square root and pseudo-inverse square root of a Hermitian positive -semi-definite matrix, from a single eigendecomposition. Equivalent -to `(sqrth_safe(M; ...), invsqrth_safe(M; ...))` but with the -eigendecomposition computed once. Eigenvalues below tolerance are clamped -to zero in both factors (Moore-Penrose convention for the inverse). - -The input must be Hermitian (as for `MatrixAlgebraKit.eigh_full`): project -with `MatrixAlgebraKit.project_hermitian` first if it is Hermitian only up -to numerical noise. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(_clamp_kwargs_doc("M")) -""" -function sqrth_invsqrth_safe(M; alg = nothing, kwargs...) - if isdiag(M) - return pow_diag_safe(M, 1 // 2; kwargs...), pow_diag_safe(M, -1 // 2; kwargs...) - end - D, V = MAK.eigh_full(M; alg) - return V * pow_diag_safe(D, 1 // 2; kwargs...) * V', - V * pow_diag_safe(D, -1 // 2; kwargs...) * V' -end - -for (gram, gram_with_pinv, eigh_full) in ( - (:gram_eigh_full, :gram_eigh_full_with_pinv, :eigh_full), - (:gram_eigh_full!!, :gram_eigh_full_with_pinv!!, :eigh_full!), - ) - @eval begin - function $gram(A::AbstractMatrix; alg = nothing, kwargs...) - D, V = MAK.$eigh_full(A; alg) - return V * sqrth_safe(D; kwargs...) - end - function $gram_with_pinv(A::AbstractMatrix; alg = nothing, kwargs...) - D, V = MAK.$eigh_full(A; alg) - return V * sqrth_safe(D; kwargs...), invsqrth_safe(D; kwargs...) * V' - end - end -end - -""" - gram_eigh_full(A::AbstractMatrix; alg=nothing, atol=0, rtol=eps(real(eltype(A)))^(3//4)) -> X - -Gram factorization of a Hermitian positive semi-definite matrix via its -eigendecomposition (balanced eigh): returns `X = V * sqrth_safe(D; atol, rtol)` -such that `A ≈ X * X'`, where `A = V * D * V'`. The square-root of `D` is -absorbed symmetrically into the two factors of the eigendecomposition. -Eigenvalues below `tol` (see [`pow_diag_safe`](@ref)) are clamped to zero. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(_clamp_kwargs_doc("A")) - -# Examples - -```jldoctest -julia> using TensorAlgebra.MatrixAlgebra: gram_eigh_full - -julia> B = [1.0 0.5; 0.5 2.0]; - -julia> A = B' * B; - -julia> X = gram_eigh_full(A); - -julia> X * X' ≈ A -true -``` - -See also [`gram_eigh_full_with_pinv`](@ref). -""" -gram_eigh_full - -""" - gram_eigh_full_with_pinv(A::AbstractMatrix; alg=nothing, atol=0, rtol=eps(real(eltype(A)))^(3//4)) -> X, Y - -Like [`gram_eigh_full`](@ref), but additionally returns -`Y = invsqrth_safe(D; atol, rtol) * V' ≈ pinv(X)`, a left inverse of `X` -on the rank subspace: `Y * X ≈ I`. Eigenvalues below `tol` are clamped to -zero in both factors. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(_clamp_kwargs_doc("A")) - -# Examples - -```jldoctest -julia> using LinearAlgebra: I - -julia> using TensorAlgebra.MatrixAlgebra: gram_eigh_full_with_pinv - -julia> B = [1.0 0.5; 0.5 2.0]; - -julia> A = B' * B; - -julia> X, Y = gram_eigh_full_with_pinv(A); - -julia> X * X' ≈ A -true - -julia> Y * X ≈ I -true -``` -""" -gram_eigh_full_with_pinv - using MatrixAlgebraKit: MatrixAlgebraKit, TruncationStrategy struct TruncationDegenerate{Strategy <: TruncationStrategy, T <: Real} <: TruncationStrategy diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 0ac1bc95..6ae6d188 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -1,9 +1,9 @@ module TensorAlgebra export contract, contract!, dual, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, - eigh_vals, gram_eigh_full, gram_eigh_full_with_pinv, invsqrth_safe, isdual, left_null, + eigh_vals, invsqrth_safe, isdual, left_null, left_orth, left_polar, lq_compact, lq_full, project_hermitian, qr_compact, - qr_full, right_null, right_orth, right_polar, sqrth_invsqrth_safe, sqrth_safe, + qr_full, right_null, right_orth, right_polar, sqrth_safe, svd_compact, svd_full, svd_trunc, svd_vals if VERSION >= v"1.11.0-DEV.469" diff --git a/src/factorizations.jl b/src/factorizations.jl index 83bb9980..87fbad56 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -63,8 +63,7 @@ end # Read-only tier: the matrix-level entries never mutate their input (they copy internally), so # the perm form consumes the maybe-alias `matricize` matricization directly. for f in ( - :gram_eigh_full, :gram_eigh_full_with_pinv, - :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, + :sqrth_safe, :invsqrth_safe, ) @eval begin function $f( @@ -90,8 +89,8 @@ for f in ( :left_polar, :right_polar, :left_orth, :right_orth, :svd_compact, :svd_full, :svd_trunc, :svd_vals, :eigh_full, :eig_full, :eigh_trunc, :eig_trunc, :eigh_vals, :eig_vals, - :left_null, :right_null, :gram_eigh_full, :gram_eigh_full_with_pinv, - :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, :project_hermitian, + :left_null, :right_null, + :sqrth_safe, :invsqrth_safe, :project_hermitian, ) @eval begin function $f(style::MatricizeStyle, A, ndims_codomain::Val{K}; kwargs...) where {K} @@ -612,124 +611,6 @@ function unmatricize_factors( return unmatricize(style, Nᴴ, (axes(Nᴴ, 1),), axes_domain) end -""" - gram_eigh_full(A, labels_A, labels_codomain, labels_domain; kwargs...) -> X - gram_eigh_full(A, perm_codomain, perm_domain; kwargs...) -> X - gram_eigh_full(A, ndims_codomain::Val; kwargs...) -> X - -Gram factorization of a generic N-dimensional array, interpreting it as a -Hermitian positive semi-definite linear map from the domain to the codomain -dimensions. Returns `X` such that `A ≈ X * X'` (contracted on the rank leg), -i.e. the codomain axes of `X` match the codomain axes of `A` and `X` has a -single trailing rank axis. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(MatrixAlgebra._clamp_kwargs_doc("A")) - -# Examples - -```jldoctest -julia> using TensorAlgebra: contract, gram_eigh_full - -julia> B = randn(3, 2, 2); - -julia> A = contract((:a, :b, :c, :d), conj(B), (:r, :a, :b), B, (:r, :c, :d)); - -julia> X = gram_eigh_full(A, (:a, :b, :c, :d), (:a, :b), (:c, :d)); - -julia> A ≈ contract((:a, :b, :c, :d), X, (:a, :b, :r), conj(X), (:c, :d, :r)) -true -``` - -See also [`gram_eigh_full_with_pinv`](@ref) and -[`MatrixAlgebra.gram_eigh_full`](@ref). -""" -gram_eigh_full - -function gram_eigh_full!!( - style::MatricizeStyle, A, ndims_codomain::Val; kwargs... - ) - A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) - X = MatrixAlgebra.gram_eigh_full!!(A_mat; kwargs...) - axes_codomain = first(bipartition(axes(A), ndims_codomain)) - return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)) -end -function gram_eigh_full!!(A, ndims_codomain::Val; kwargs...) - return gram_eigh_full!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) -end - -function unmatricize_factors( - ::typeof(gram_eigh_full), style::MatricizeStyle, X, - axes_codomain, axes_domain - ) - return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)) -end - -""" - gram_eigh_full_with_pinv(A, labels_A, labels_codomain, labels_domain; kwargs...) -> X, Y - gram_eigh_full_with_pinv(A, perm_codomain, perm_domain; kwargs...) -> X, Y - gram_eigh_full_with_pinv(A, ndims_codomain::Val; kwargs...) -> X, Y - -Like [`gram_eigh_full`](@ref), but additionally returns `Y ≈ pinv(X)` such -that `Y * X ≈ I` on the rank subspace (a left inverse). The codomain axes -of `X` match the codomain axes of `A`; `Y` has a leading rank axis followed -by the codomain axes. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(MatrixAlgebra._clamp_kwargs_doc("A")) - -# Examples - -```jldoctest -julia> using LinearAlgebra: I - -julia> using TensorAlgebra: contract, gram_eigh_full_with_pinv - -julia> B = randn(8, 2, 2); - -julia> A = contract((:a, :b, :c, :d), conj(B), (:r, :a, :b), B, (:r, :c, :d)); - -julia> X, Y = gram_eigh_full_with_pinv(A, (:a, :b, :c, :d), (:a, :b), (:c, :d)); - -julia> A ≈ contract((:a, :b, :c, :d), X, (:a, :b, :r), conj(X), (:c, :d, :r)) -true - -julia> contract((:r, :s), Y, (:r, :a, :b), X, (:a, :b, :s)) ≈ I -true -``` - -See also [`MatrixAlgebra.gram_eigh_full_with_pinv`](@ref). -""" -gram_eigh_full_with_pinv - -function gram_eigh_full_with_pinv!!( - style::MatricizeStyle, A, ndims_codomain::Val; kwargs... - ) - A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) - X, Y = MatrixAlgebra.gram_eigh_full_with_pinv!!(A_mat; kwargs...) - axes_codomain = first(bipartition(axes(A), ndims_codomain)) - return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)), - unmatricize(style, Y, (axes(Y, 1),), axes_codomain) -end -function gram_eigh_full_with_pinv!!(A, ndims_codomain::Val; kwargs...) - return gram_eigh_full_with_pinv!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) -end - -function unmatricize_factors( - ::typeof(gram_eigh_full_with_pinv), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) - X, Y = F - return unmatricize(style, X, axes_codomain, (conj(axes(X, ndims(X))),)), - unmatricize(style, Y, (axes(Y, 1),), axes_codomain) -end - """ sqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P sqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P @@ -748,7 +629,7 @@ up to numerical noise. $(MatrixAlgebra._clamp_kwargs_doc("A")) -See also [`invsqrth_safe`](@ref), [`sqrth_invsqrth_safe`](@ref), and +See also [`invsqrth_safe`](@ref) and [`MatrixAlgebra.sqrth_safe`](@ref). """ sqrth_safe @@ -771,7 +652,7 @@ first if it is Hermitian only up to numerical noise. $(MatrixAlgebra._clamp_kwargs_doc("A")) -See also [`sqrth_safe`](@ref), [`sqrth_invsqrth_safe`](@ref), and +See also [`sqrth_safe`](@ref) and [`MatrixAlgebra.invsqrth_safe`](@ref). """ invsqrth_safe @@ -807,35 +688,6 @@ function unmatricize_factors( return unmatricize(style, H_mat, axes_codomain, axes_domain) end -""" - sqrth_invsqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P, Pinv - sqrth_invsqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P, Pinv - sqrth_invsqrth_safe(A, ndims_codomain::Val; kwargs...) -> P, Pinv - -Square root and pseudo-inverse square root of a generic N-dimensional -array (see [`sqrth_safe`](@ref) and [`invsqrth_safe`](@ref)), from a -single eigendecomposition. Both results carry the same codomain and -domain axes as `A`. - -## Keyword arguments - - - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. - -$(MatrixAlgebra._clamp_kwargs_doc("A")) - -See also [`MatrixAlgebra.sqrth_invsqrth_safe`](@ref). -""" -sqrth_invsqrth_safe - -function unmatricize_factors( - ::typeof(sqrth_invsqrth_safe), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) - P_mat, Pinv_mat = F - return unmatricize(style, P_mat, axes_codomain, axes_domain), - unmatricize(style, Pinv_mat, axes_codomain, axes_domain) -end - """ TensorAlgebra.one(A, labels_A, labels_codomain, labels_domain) -> Id TensorAlgebra.one(A, perm_codomain, perm_domain) -> Id diff --git a/test/test_exports.jl b/test/test_exports.jl index 0ff6b33e..5cb1aea4 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -13,8 +13,6 @@ using Test: @test, @testset :eigh_full, :eigh_trunc, :eigh_vals, - :gram_eigh_full, - :gram_eigh_full_with_pinv, :invsqrth_safe, :isdual, :left_null, @@ -28,7 +26,6 @@ using Test: @test, @testset :right_null, :right_orth, :right_polar, - :sqrth_invsqrth_safe, :sqrth_safe, :svd_compact, :svd_full, @@ -56,15 +53,12 @@ using Test: @test, @testset exports = [ :MatrixAlgebra, - :gram_eigh_full, - :gram_eigh_full_with_pinv, :invsqrt_diag_safe, :invsqrth_safe, :pow_diag_safe, :pow_diag_safe!, :powh_safe, :sqrt_diag_safe, - :sqrth_invsqrth_safe, :sqrth_safe, ] @test issetequal(names(TensorAlgebra.MatrixAlgebra), exports) diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index d69a007d..d9f62c63 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -1,9 +1,8 @@ using LinearAlgebra: LinearAlgebra, Diagonal, I, diag, norm using MatrixAlgebraKit: truncrank using TensorAlgebra: TensorAlgebra, contract, eig_full, eig_vals, eigh_full, eigh_vals, - gram_eigh_full, gram_eigh_full_with_pinv, left_null, left_orth, left_polar, lq_compact, - lq_full, qr_compact, qr_full, right_null, right_orth, right_polar, svd_compact, - svd_full, svd_trunc, svd_vals + left_null, left_orth, left_polar, lq_compact, lq_full, qr_compact, qr_full, right_null, + right_orth, right_polar, svd_compact, svd_full, svd_trunc, svd_vals using Test: @test, @testset using TestExtras: @constinferred @@ -306,59 +305,6 @@ end @test A ≈ contract(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) end -# Gram factorization -# ------------------ -# Build a Hermitian positive semi-definite tensor A[a,b,c,d] with codomain -# (a, b) and domain (c, d): pick a random B[k, a, b] (k = aux), then form -# A = B' * B over k. By construction A ≈ X' * X for X[r, a, b] with rank r -# bounded by k (rank leg first, following the Cholesky `A = U' * U` -# convention). -@testset "Full-rank gram_eigh_full ($T)" for T in elts - B = randn(T, 6, 2, 3) # k = 6, codomain = (a, b) of size 2*3 = 6 -> full rank - A = contract((:a, :b, :c, :d), conj(B), (:k, :a, :b), B, (:k, :c, :d)) - labels_A = (:a, :b, :c, :d) - labels_X = (:a, :b) - labels_Y = (:c, :d) - - Acopy = copy(A) - X = @constinferred gram_eigh_full(A, labels_A, labels_X, labels_Y) - @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, X, (:a, :b, :r), conj(X), (:c, :d, :r)) - @test A ≈ A′ - @test size(X, ndims(X)) == size(A, 1) * size(A, 2) - - # `Val`, perm, and label entries agree. - @test gram_eigh_full(A, Val(2)) ≈ X - @test gram_eigh_full(A, (1, 2), (3, 4)) ≈ X - - # `with_pinv` variant: Y is a left inverse of X (Y * X ≈ I on the - # rank subspace). - X2, Y2 = @constinferred gram_eigh_full_with_pinv(A, labels_A, labels_X, labels_Y) - @test A ≈ contract(labels_A, X2, (:a, :b, :r), conj(X2), (:c, :d, :r)) - YX = contract((:r, :s), Y2, (:r, :a, :b), X2, (:a, :b, :s)) - @test YX ≈ I -end - -@testset "Rank-deficient gram_eigh_full ($T)" for T in elts - B = randn(T, 4, 2, 3) # k = 4 < codomain dim 6, so A is rank-4 - A = contract((:a, :b, :c, :d), conj(B), (:k, :a, :b), B, (:k, :c, :d)) - - # Recovery of A is independent of the `rtol` cutoff because all - # nonzero eigenvalues sit far above any reasonable threshold. - X = gram_eigh_full(A, Val(2); rtol = 1.0e-10) - @test A ≈ contract( - (:a, :b, :c, :d), X, (:a, :b, :r), conj(X), (:c, :d, :r) - ) - - # Moore–Penrose-like identity: X * Y * X ≈ X when Y is pinv(X). With - # cod-first X and rank-first Y, contract Y[r, a, b] * X[a, b, s] → P[r, s] - # (projector onto the rank subspace), then X * P → X. - X2, Y2 = gram_eigh_full_with_pinv(A, Val(2); rtol = 1.0e-10) - P = contract((:r, :s), Y2, (:r, :a, :b), X2, (:a, :b, :s)) - XP = contract((:c, :d, :r), X2, (:c, :d, :s), P, (:s, :r)) - @test XP ≈ X2 -end - # one (identity tensor) # --------------------- # An identity tensor matricized along its codomain/domain partition is the diff --git a/test/test_matrixalgebra.jl b/test/test_matrixalgebra.jl index e279ea4d..bda1a277 100644 --- a/test/test_matrixalgebra.jl +++ b/test/test_matrixalgebra.jl @@ -147,42 +147,6 @@ elts = (Float32, Float64, ComplexF32, ComplexF64) @test norm(ũ * s̃ * ṽ) ≈ 0 end - @testset "gram_eigh_full" begin - n = 5 - # Full-rank Hermitian PSD. Use a tall random factor so `B' * B` - # is comfortably full rank even at Float32 precision (a square - # random `B` can produce a `B' * B` whose smallest eigenvalue - # falls below the default rtol clamp on some seeds). - rng = StableRNG(123) - B = randn(rng, elt, 2n, n) - A = B' * B - X = MatrixAlgebra.gram_eigh_full(A) - @test X * X' ≈ A - @test size(X) == (n, n) - - X2, Y2 = MatrixAlgebra.gram_eigh_full_with_pinv(A) - @test X2 * X2' ≈ A - @test Y2 * X2 ≈ I(n) - - # `!!` variant accepts a destroyable copy. - Xb = MatrixAlgebra.gram_eigh_full!!(copy(A)) - @test Xb * Xb' ≈ A - - # Rank deficient: A is n×n of rank k < n. Recovery of A still holds; - # X * Y is the projector onto the rank-k codomain subspace - # (idempotent, rank k), and X * P ≈ X (Moore–Penrose). - k = 3 - Brd = randn(rng, elt, k, n) - Ard = Brd' * Brd - Xrd, Yrd = MatrixAlgebra.gram_eigh_full_with_pinv( - Ard; rtol = sqrt(eps(real(elt))) - ) - @test Xrd * Xrd' ≈ Ard - P = Xrd * Yrd - @test P * P ≈ P - @test P * Xrd ≈ Xrd - end - @testset "powh_safe / sqrth_safe / invsqrth_safe" begin n = 4 rng = StableRNG(123) From ec471e83f2224da38c5d73feab8cdffd5ce6c7dd Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 15 Sep 2026 22:56:31 -0400 Subject: [PATCH 14/33] Let a style overload the total matricize copy A graded array already gets owned storage out of `permutedimsop`, whose stored matrix is the answer, so decomposing the copy into an allocation plus an in-place write would copy that storage a second time. Co-Authored-By: Claude Opus 5 (1M context) --- src/matricize.jl | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/src/matricize.jl b/src/matricize.jl index 23cff0f2..830872ee 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -86,6 +86,12 @@ end # needs the fused axes, which only the style knows: TensorAlgebra deliberately has no generic # axis-fusion interface. It is also what makes the copy path terminate, since `matricizeop!` is a # distinct function from the router rather than a re-entry into it. +# +# `matricizeopcopy` is itself an overload point for a style whose owned matricization already falls +# out of an allocating operation it has (for a graded array, permuting into fresh storage whose +# stored matrix is the answer). Such a style overloads the copy instead of +# `allocate_output`/`matricizeop!`, which would copy that storage a second time, and then owes +# only `matricizeopview` and `is_output_view`. """ matricizeop(op, a, perm_codomain, perm_domain) From 4a830205e125fb9ffbf94f2bf3f9dedbc1e52abc Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 16 Sep 2026 00:17:59 -0400 Subject: [PATCH 15/33] Give the contract interface one name per argument form The labels forms keep the plain names and the bipermutation forms take a `perm` marker, freeing `contract` to become variadic over operands. `allocate_contract_output` is gone in favor of overloading `allocate_output`, and a generic `select_algorithm` sits above the per-operation resolvers. Co-Authored-By: Claude Opus 5 (1M context) --- ext/TensorAlgebraTensorKitExt.jl | 2 +- .../TensorAlgebraTensorOperationsExt.jl | 8 +- src/TensorAlgebra.jl | 6 +- src/algorithm.jl | 39 +++++ src/contract/allocate_output.jl | 18 +-- src/contract/contract.jl | 134 ++++++++++++++---- src/contract/contract_matricize.jl | 2 +- src/diagonal.jl | 34 +++-- src/factorizations.jl | 6 +- test/test_basics.jl | 41 +++--- test/test_exports.jl | 10 +- test/test_factorizations.jl | 82 ++++++----- test/test_matricize.jl | 2 +- test/test_mooncakeext.jl | 6 +- 14 files changed, 279 insertions(+), 111 deletions(-) create mode 100644 src/algorithm.jl diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 58fa7093..90bbb924 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -293,7 +293,7 @@ end # already implements through its TensorOperations interface. Route the generic `contract` # there: `zero!` clears the `similar_map`-allocated destination, and the default algorithm # hands the in-place contraction to the TensorOperations backend (see the TensorOperations -# extension's `contractopadd!`). +# extension's `contractpermopadd!`). TensorAlgebra.zero!(t::AbstractTensorMap) = VectorInterface.zerovector!(t) # A `TensorMap` is not an `AbstractArray`, so the generic in-place `TensorAlgebra` operations diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 2b62f80e..79c84f3c 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -22,7 +22,7 @@ end # ---------------------------------------------------------------- # not in-place -function TA.contract( +function TA.contractperm( algorithm::TensorOperationsAlgorithm, perm_dest_codomain, perm_dest_domain, a1::AbstractArray, perm1_codomain, perm1_domain, @@ -39,7 +39,7 @@ function TA.contract( ) end -function TA.contract( +function TA.contractalign( algorithm::TensorOperationsAlgorithm, labels_dest, a1::AbstractArray, labels1, @@ -56,7 +56,7 @@ function TA.contract( end # in-place -function TA.contractopadd!( +function TA.contractpermopadd!( algorithm::TensorOperationsAlgorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, @@ -89,7 +89,7 @@ function TO.tensorcontract!( ) op1 = conj1 ? conj : identity op2 = conj2 ? conj : identity - return TA.contractopadd!( + return TA.contractpermopadd!( backend, a_dest, permblocks_dest..., op1, a1, permblocks1..., diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 6ae6d188..b3b04fd7 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -1,6 +1,7 @@ module TensorAlgebra -export contract, contract!, dual, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, +export contract, contract!, contractalign, dual, eig_full, eig_trunc, eig_vals, eigh_full, + eigh_trunc, eigh_vals, invsqrth_safe, isdual, left_null, left_orth, left_polar, lq_compact, lq_full, project_hermitian, qr_compact, qr_full, right_null, right_orth, right_polar, sqrth_safe, @@ -9,7 +10,7 @@ export contract, contract!, dual, eig_full, eig_trunc, eig_vals, eigh_full, eigh if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public biperm, bipartition, cat_similar, concatenate, concatenate!, ContractAlgorithm, contractopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" + "public allocate_output, biperm, bipartition, cat_similar, check_input, concatenate, concatenate!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, output_axes, select_algorithm, default_algorithm, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" ) ) end @@ -25,6 +26,7 @@ include("concatenate.jl") include("directsum.jl") include("dual.jl") include("to_range.jl") +include("algorithm.jl") include("contract/contractalgorithm.jl") include("contract/contract.jl") include("contract/contract_labels.jl") diff --git a/src/algorithm.jl b/src/algorithm.jl new file mode 100644 index 00000000..28775bb8 --- /dev/null +++ b/src/algorithm.jl @@ -0,0 +1,39 @@ +""" + TensorAlgebra.select_algorithm(f, args...; alg = nothing, kwargs...) + +Resolve the algorithm operation `f` should run with on `args`. An `alg` of `nothing` defers to +[`TensorAlgebra.default_algorithm`](@ref); anything else is validated and passed through. + +This is the forward-facing layer the per-operation resolvers sit under, so a caller that does not +care which operation it is dispatching writes `select_algorithm(f, ...)` and a backend registers +its choice on `default_algorithm(f, ...)`. +""" +function select_algorithm(f, args...; alg = nothing, kwargs...) + isnothing(alg) && return default_algorithm(f, args...; kwargs...) + return select_algorithm_specified(f, alg, args...; kwargs...) +end + +""" + TensorAlgebra.default_algorithm(f, args...) + TensorAlgebra.default_algorithm(f, argtypes::Type...) + +The algorithm operation `f` runs with on `args` when the caller names none. The types form is the +registration point for a storage type; the values form defaults to it. + +Each operation bridges to its own resolver, so `default_algorithm(contract, A1, A2)` is +[`TensorAlgebra.default_contract_algorithm`](@ref). +""" +function default_algorithm(f, args...; kwargs...) + return default_algorithm(f, map(typeof, args)...; kwargs...) +end +function default_algorithm(f, argtypes::Type...; kwargs...) + return throw(MethodError(default_algorithm, (f, argtypes...))) +end + +# `alg` named something. A resolved algorithm object passes through; anything else is a caller +# error, reported against the operation rather than as a `MethodError` from inside the resolver. +function select_algorithm_specified(f, alg, args...; kwargs...) + return throw( + ArgumentError("`$alg` is not an algorithm for `$f`") + ) +end diff --git a/src/contract/allocate_output.jl b/src/contract/allocate_output.jl index 13b392c0..5fb32103 100644 --- a/src/contract/allocate_output.jl +++ b/src/contract/allocate_output.jl @@ -68,13 +68,16 @@ function output_axes( axes_uncontracted, perm_dest_codomain, perm_dest_domain ) # The operand axes are stored/dualized, so un-dualize the domain axes into the codomain-facing - # construction convention shared by `allocate_contract_output`, `similar_map`, and `unmatricize` - # (a no-op on dense axes). + # construction convention shared by `similar_map` and `unmatricize` (a no-op on dense axes). return axes_codomain_dest, conj.(axes_domain_dest) end # TODO: Use `ArrayLayouts`-like `MulAdd` object, # i.e. `ContractAdd`? +# The destination `contract` writes into. A structured operand type overloads this directly, deriving +# the axes and element type from `output_axes` and `Base.promote_op` as below; the permutations are +# part of the signature because the contraction pattern is not recoverable from the destination leg +# counts alone. function allocate_output( ::typeof(contract), perm_dest_codomain, perm_dest_domain, @@ -97,14 +100,5 @@ function allocate_output( a2, perm2_codomain, perm2_domain ) T = Base.promote_op(matprod, eltype(a1), eltype(a2)) - return allocate_contract_output(a1, a2, T, axes_codomain_dest, axes_domain_dest) -end - -# Allocate the output container for `contract`: the operand types, the output element type and -# axes (domain codomain-facing), and the output's codomain/domain leg counts (the axes tuple -# lengths) select the container type. Internal to TensorAlgebra, not a public extension point: -# the leg counts identify the contraction pattern only for matrix-shaped operands (see the -# `Diagonal` method in `diagonal.jl`), so external structured types should not overload it. -function allocate_contract_output(a1, a2, T, axes_codomain::Tuple, axes_domain::Tuple) - return zero!(similar_map(a1, T, axes_codomain, axes_domain)) + return zero!(similar_map(a1, T, axes_codomain_dest, axes_domain_dest)) end diff --git a/src/contract/contract.jl b/src/contract/contract.jl index f31f14ea..a3d47e17 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -2,15 +2,65 @@ # TODO: Add `scaledcontract(a1, labels1, a2, labels2, α) = α * contract(a1, labels1, a2, labels2)`. # contract (labels) +""" + contract(a1, labels1, a2, labels2, ...; alg = nothing) -> a_dest, labels_dest + +Contract the arrays over the labels they share, returning the result along with the labels of its +dimensions. A label appearing on two operands is summed over, one appearing on a single operand +survives, and `labels_dest` reports the surviving labels in the order the result carries them. + +Operands past the second are contracted one pair at a time from left to right, so the call +expresses the contraction order rather than requesting an optimized one. + +```jldoctest +julia> using TensorAlgebra: contract + +julia> a, b = randn(2, 3), randn(3, 4); + +julia> ab, labels = contract(a, (:i, :j), b, (:j, :k)); + +julia> labels +(:i, :k) + +julia> ab ≈ a * b +true +``` + +See also [`contractalign`](@ref) to name the output dimensions and get back just the array. +""" function contract(a1, labels1, a2, labels2; kwargs...) # Optionally convert the labels to a representation cheaper to run the bookkeeping on (see # `label_type`). `encode_contraction_labels`/`decode_contraction_labels` are no-ops unless the label type opts in. l1, l2 = encode_contraction_labels(labels1, labels2) l_dest = contract_labels(l1, l2) - a_dest = contract(l_dest, a1, l1, a2, l2; kwargs...) + a_dest = contractalign(l_dest, a1, l1, a2, l2; kwargs...) return a_dest, decode_contraction_labels(l_dest, labels1, labels2) end -function contract( +function contract(a1, labels1, a2, labels2, a3, labels3, rest...; kwargs...) + check_alternating_labels(contract, rest) + a12, labels12 = contract(a1, labels1, a2, labels2; kwargs...) + return contract(a12, labels12, a3, labels3, rest...; kwargs...) +end + +""" + contractalign(labels_dest, a1, labels1, a2, labels2, ...; alg = nothing) -> a_dest + +Contract the arrays over the labels they share into a result whose dimensions carry +`labels_dest`, which must be the surviving labels in some order. + +This is [`contract`](@ref) with the output specified, so it returns the array on its own. The +name matches `ITensorBase.align`: arrange the result's dimensions to match the labels given. + +```jldoctest +julia> using TensorAlgebra: contractalign + +julia> a, b = randn(2, 3), randn(3, 4); + +julia> contractalign((:k, :i), a, (:i, :j), b, (:j, :k)) ≈ permutedims(a * b, (2, 1)) +true +``` +""" +function contractalign( labels_dest, a1, labels1, a2, labels2; kwargs... ) t1 = ntuple(i -> labels1[i], Val(ndims(a1))) @@ -18,7 +68,7 @@ function contract( contracted1 = map(in(t2), t1) # Cross into a `Val(K)` method (a function-barrier on the contracted count) so the # bipartitioned permutations and the contraction below them are type-stable. - return _contract( + return _contractalign( Val(count(contracted1)), labels_dest, a1, @@ -29,17 +79,38 @@ function contract( kwargs... ) end -function _contract( +# Only the last pair lands on the requested output labels; the ones before it infer their own. +function contractalign( + labels_dest, a1, labels1, a2, labels2, a3, labels3, rest...; kwargs... + ) + check_alternating_labels(contractalign, rest) + a12, labels12 = contract(a1, labels1, a2, labels2; kwargs...) + return contractalign(labels_dest, a12, labels12, a3, labels3, rest...; kwargs...) +end +function _contractalign( ::Val{K}, labels_dest, a1, labels1, a2, labels2, contracted1; kwargs... ) where {K} biperm_dest, biperm1, biperm2 = biperms(contract, Val(K), labels_dest, labels1, labels2, contracted1) - return contract(biperm_dest..., a1, biperm1..., a2, biperm2...; kwargs...) + return contractperm(biperm_dest..., a1, biperm1..., a2, biperm2...; kwargs...) end -# contract (bipartitioned permutations) -function contract( +# The variadic forms take arrays and labels in alternating positions, so a trailing group with an +# odd length is a miscount at the call site rather than something to diagnose further down. +function check_alternating_labels(f, rest::Tuple) + iseven(length(rest)) || throw( + ArgumentError( + "`$f` takes each array followed by its labels, so the trailing arguments must come in pairs" + ) + ) + return nothing +end + +# contractperm (bipartitioned permutations) +# `perm` marks the whole biperm ladder: every rung has a labels-form sibling under the plain name, +# and once `contract` is variadic over operands the two can no longer be told apart by arity. +function contractperm( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; kwargs... @@ -48,14 +119,14 @@ function contract( Ndest = Val(length(perm1_codomain) + length(perm2_domain)) perm_dest_codomain, perm_dest_domain = bipartition(ntuple(identity, Ndest), Ndest_codomain) - return contract( + return contractperm( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; kwargs... ) end -function contract( +function contractperm( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; @@ -67,7 +138,7 @@ function contract( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) - return contract!( + return contractperm!( a_dest, perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; @@ -86,13 +157,13 @@ function contract!( a_dest, labels_dest, a1, labels1, a2, labels2, true, false; kwargs... ) end -function contract!( +function contractperm!( a_dest, perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; kwargs... ) - return contractadd!( + return contractpermadd!( a_dest, perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain, @@ -112,15 +183,15 @@ function contractadd!( a_dest, labels_dest, identity, a1, labels1, identity, a2, labels2, α, β; kwargs... ) end -# contractadd! (bipartitioned permutations) -function contractadd!( +# contractpermadd! (bipartitioned permutations) +function contractpermadd!( a_dest, perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain, α::Number, β::Number; kwargs... ) - return contractopadd!( + return contractpermopadd!( a_dest, perm_dest_codomain, perm_dest_domain, identity, a1, perm1_codomain, perm1_domain, identity, a2, perm2_codomain, perm2_domain, @@ -153,17 +224,17 @@ function _contractopadd!( ) where {K} biperm_dest, biperm1, biperm2 = biperms(contract, Val(K), labels_dest, labels1, labels2, contracted1) - return contractopadd!( + return contractpermopadd!( a_dest, biperm_dest..., op1, a1, biperm1..., op2, a2, biperm2..., α, β; kwargs... ) end -# contractopadd! (bipartitioned permutations, algorithm selection) -function contractopadd!( +# contractpermopadd! (bipartitioned permutations, algorithm selection) +function contractpermopadd!( a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, α::Number, β::Number; - alg = DefaultContractAlgorithm(), kwargs... + alg = nothing, kwargs... ) check_input( contract!, @@ -171,8 +242,8 @@ function contractopadd!( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) - algorithm = select_contract_algorithm(alg, a1, a2; kwargs...) - return contractopadd!( + algorithm = select_algorithm(contract, a1, a2; alg, kwargs...) + return contractpermopadd!( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, @@ -180,9 +251,9 @@ function contractopadd!( α, β ) end -# contractopadd! (dispatched on the algorithm, bipartitioned permutations) +# contractpermopadd! (dispatched on the algorithm, bipartitioned permutations) # Required interface if not using matricized contraction -function contractopadd!( +function contractpermopadd!( algorithm::ContractAlgorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, @@ -191,7 +262,7 @@ function contractopadd!( ) return throw( MethodError( - contractopadd!, + contractpermopadd!, ( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, @@ -202,3 +273,18 @@ function contractopadd!( ) ) end + +# Bridges from the operation-generic algorithm layer in `algorithm.jl` down to the contraction +# resolvers. They live here rather than beside those resolvers because dispatching on +# `::typeof(contract)` needs `contract` to exist, and `contractalgorithm.jl` is included first for +# the algorithm types this file's signatures use. +function default_algorithm(::typeof(contract), A1::Type, A2::Type; kwargs...) + algorithm = default_contract_algorithm(A1, A2) + reject_algorithm_kwargs(algorithm; kwargs...) + return algorithm +end +function select_algorithm_specified( + ::typeof(contract), alg::ContractAlgorithm, a1, a2; kwargs... + ) + return select_contract_algorithm(alg, a1, a2; kwargs...) +end diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 0b5fc951..409d032d 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -1,6 +1,6 @@ using LinearAlgebra: mul! -function contractopadd!( +function contractpermopadd!( algorithm::Matricize, a_dest::AbstractArray, biperm_dest_codomain, biperm_dest_domain, op1, a1::AbstractArray, biperm1_codomain, biperm1_domain, diff --git a/src/diagonal.jl b/src/diagonal.jl index 0488effa..d4befc6e 100644 --- a/src/diagonal.jl +++ b/src/diagonal.jl @@ -69,13 +69,31 @@ function unmatricize( return unmatricize(style, copyto!(similar(m, axes(m)), m), axes_codomain, axes_domain) end -# Contracting two `Diagonal`s over a single leg is the matmul/endomorphism pattern -# `Diagonal * Diagonal = Diagonal` (all transpose variants `[i,j]*[j,k]`, `[i,j]*[k,j]`, ...), -# whose `{1,1}` output stays `Diagonal`, so allocate one. Every other output shape (rank-4 outer -# product, scalar full contraction) is not representable as a `Diagonal` and falls back to the -# generic dense allocation, matching `Diagonal`/dense mixing. -function allocate_contract_output( - a1::Diagonal, a2::Diagonal, T, axes_codomain::Tuple{Any}, axes_domain::Tuple{Any} +function allocate_output( + ::typeof(contract), + perm_dest_codomain, perm_dest_domain, + a1::Diagonal, perm1_codomain, perm1_domain, + a2::Diagonal, perm2_codomain, perm2_domain + ) + check_input( + contract, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain + ) + axes_codomain_dest, axes_domain_dest = output_axes( + contract, + perm_dest_codomain, perm_dest_domain, + a1, perm1_codomain, perm1_domain, + a2, perm2_codomain, perm2_domain ) - return Diagonal(zero!(similar(a1.diag, T, (only(axes_codomain),)))) + T = Base.promote_op(matprod, eltype(a1), eltype(a2)) + # Contracting two `Diagonal`s over a single leg, leaving one free leg on each, is the + # matmul/endomorphism pattern `Diagonal * Diagonal = Diagonal` (all transpose variants + # `[i,j]*[j,k]`, `[i,j]*[k,j]`, ...), whose `{1,1}` output stays `Diagonal`. Every other + # pattern (rank-4 outer product, scalar full contraction) is not representable as a + # `Diagonal` and takes the generic dense allocation, matching `Diagonal`/dense mixing. + is_matmul = + length(perm1_codomain) == 1 && length(perm1_domain) == 1 && + length(perm2_domain) == 1 + is_matmul || + return zero!(similar_map(a1, T, axes_codomain_dest, axes_domain_dest)) + return Diagonal(zero!(similar(a1.diag, T, (only(axes_codomain_dest),)))) end diff --git a/src/factorizations.jl b/src/factorizations.jl index 87fbad56..63440716 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -429,15 +429,15 @@ truncation error `ϵ`, the 2-norm of the discarded singular values. # Examples ```jldoctest -julia> using TensorAlgebra: svd_trunc, contract +julia> using TensorAlgebra: svd_trunc, contractalign julia> A = randn(4, 4); julia> U, S, Vᴴ, ϵ = svd_trunc(A, (:i, :j), (:i,), (:j,)); -julia> SV = contract((:u, :j), S, (:u, :v), Vᴴ, (:v, :j)); +julia> SV = contractalign((:u, :j), S, (:u, :v), Vᴴ, (:v, :j)); -julia> contract((:i, :j), U, (:i, :u), SV, (:u, :j)) ≈ A +julia> contractalign((:i, :j), U, (:i, :u), SV, (:u, :j)) ≈ A true julia> isapprox(ϵ, 0; atol = 1e-10) diff --git a/test/test_basics.jl b/test/test_basics.jl index b0a119cb..2ffabf85 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,8 +1,8 @@ import TensorAlgebra using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, - contract!, contractadd!, length_codomain, length_domain, matricize, unmatricize, - unmatricize! + contract!, contractadd!, contractalign, length_codomain, length_domain, matricize, + unmatricize, unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -172,15 +172,22 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int a2 = randn(3, 4) # A keyword no algorithm can consume must name itself, not surface as a `MethodError` # from inside the resolver. - @test_throws ArgumentError contract((1, 3), a1, (1, 2), a2, (2, 3); nonsense = 1) - @test_throws ArgumentError contract( + @test_throws ArgumentError contractalign( + (1, 3), + a1, + (1, 2), + a2, + (2, 3); + nonsense = 1 + ) + @test_throws ArgumentError contractalign( (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize(), nonsense = 1 ) # A non-algorithm passed as `alg` says so rather than erroring with "Not implemented". @test_throws ArgumentError TensorAlgebra.select_contract_algorithm(:nope, a1, a2) # The supported spellings still work. - @test contract((1, 3), a1, (1, 2), a2, (2, 3)) ≈ a1 * a2 - @test contract( + @test contractalign((1, 3), a1, (1, 2), a2, (2, 3)) ≈ a1 * a2 + @test contractalign( (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize() ) ≈ a1 * a2 end @@ -201,8 +208,8 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int a_dest = ones(elt_dest, (1, 1)) @test_throws ArgumentError contract(a1, (1, 2, 4), a2, (2, 3)) @test_throws ArgumentError contract(a1, (1, 2), a2, (2, 3, 4)) - @test_throws ArgumentError contract((1, 3, 4), a1, (1, 2), a2, (2, 3)) - @test_throws ArgumentError contract((1, 3), a1, (1, 2), a2, (2, 4)) + @test_throws ArgumentError contractalign((1, 3, 4), a1, (1, 2), a2, (2, 3)) + @test_throws ArgumentError contractalign((1, 3), a1, (1, 2), a2, (2, 4)) @test_throws ArgumentError contract!(a_dest, (1, 3, 4), a1, (1, 2), a2, (2, 3)) dims = (2, 3, 4, 5, 6, 7, 8, 9, 10) @@ -241,14 +248,14 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test a_dest ≈ a_dest_tensoroperations # Specify destination labels - a_dest = contract(labels_dest, a1, labels1, a2, labels2) - a_dest_tensoroperations = contract( + a_dest = contractalign(labels_dest, a1, labels1, a2, labels2) + a_dest_tensoroperations = contractalign( labels_dest, a1, labels1, a2, labels2; alg = alg_tensoroperations ) @test a_dest ≈ a_dest_tensoroperations - a_dest = contract(labels_dest′, a1, labels1, a2, labels2) - a_dest_tensoroperations = contract( + a_dest = contractalign(labels_dest′, a1, labels1, a2, labels2) + a_dest_tensoroperations = contractalign( labels_dest′, a1, labels1, a2, labels2; alg = alg_tensoroperations ) @test a_dest ≈ a_dest_tensoroperations @@ -286,7 +293,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test labels_dest == [L(1), L(4)] # Specifying the destination labels still works for opted-in types. - a_dest = contract([L(1), L(4)], a1, (L(1), L(2), L(3)), a2, (L(2), L(3), L(4))) + a_dest = contractalign([L(1), L(4)], a1, (L(1), L(2), L(3)), a2, (L(2), L(3), L(4))) @test a_dest ≈ a_ref # Empty labels (e.g. a scalar operand) are handled. @@ -309,7 +316,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test eltype(a_dest) === elt_dest @test a_dest ≈ reshape(vec(a1) * transpose(vec(a2)), (size(a1)..., size(a2)...)) - a_dest = contract(("i", "k", "j", "l"), a1, ("i", "j"), a2, ("k", "l")) + a_dest = contractalign(("i", "k", "j", "l"), a1, ("i", "j"), a2, ("k", "l")) @test eltype(a_dest) === elt_dest @test a_dest ≈ permutedims( reshape(vec(a1) * transpose(vec(a2)), (size(a1)..., size(a2)...)), (1, 3, 2, 4) @@ -456,17 +463,17 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int size_dest_example = (3, 5, 2, 4) # Array-scalar contraction. - a_dest = contract(labels_dest_example, a, labels_a, s, ()) + a_dest = contractalign(labels_dest_example, a, labels_a, s, ()) @test size(a_dest) == size_dest_example @test a_dest ≈ permutedims(a, (2, 4, 1, 3)) * s[] # Scalar-array contraction. - a_dest = contract(labels_dest_example, s, (), a, labels_a) + a_dest = contractalign(labels_dest_example, s, (), a, labels_a) @test size(a_dest) == size_dest_example @test a_dest ≈ permutedims(a, (2, 4, 1, 3)) * s[] # Scalar-scalar contraction. - a_dest = contract((), s, (), t, ()) + a_dest = contractalign((), s, (), t, ()) @test size(a_dest) == () @test a_dest[] ≈ s[] * t[] diff --git a/test/test_exports.jl b/test/test_exports.jl index 5cb1aea4..89f67df8 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -6,6 +6,7 @@ using Test: @test, @testset :TensorAlgebra, :contract, :contract!, + :contractalign, :dual, :eig_full, :eig_trunc, @@ -37,12 +38,15 @@ using Test: @test, @testset append!( exports, [ - :biperm, :bipartition, :cat_similar, - :concatenate, :concatenate!, :ContractAlgorithm, :contractopadd!, :data, + :allocate_output, :biperm, :bipartition, :cat_similar, :check_input, + :concatenate, :concatenate!, :ContractAlgorithm, :contractopadd!, + :contractperm, :contractperm!, :contractpermadd!, :contractpermopadd!, + :data, :datatype, :directsum, :flattenlinear, :is_output_view, :label_type, :matricize, :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, - :permutedims, :permutedims!, :scalar, :similar_map, + :default_algorithm, :output_axes, :permutedims, :select_algorithm, + :permutedims!, :scalar, :similar_map, :TensorOperationsAlgorithm, :to_range, :tr, :tryflattenlinear, :ungrade, :zero!, :scale!, :permuteddims, :PermutedDims, diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index d9f62c63..c2daeaf2 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -1,8 +1,8 @@ using LinearAlgebra: LinearAlgebra, Diagonal, I, diag, norm using MatrixAlgebraKit: truncrank -using TensorAlgebra: TensorAlgebra, contract, eig_full, eig_vals, eigh_full, eigh_vals, - left_null, left_orth, left_polar, lq_compact, lq_full, qr_compact, qr_full, right_null, - right_orth, right_polar, svd_compact, svd_full, svd_trunc, svd_vals +using TensorAlgebra: TensorAlgebra, contract, contractalign, eig_full, eig_vals, eigh_full, + eigh_vals, left_null, left_orth, left_polar, lq_compact, lq_full, qr_compact, qr_full, + right_null, right_orth, right_polar, svd_compact, svd_full, svd_trunc, svd_vals using Test: @test, @testset using TestExtras: @constinferred @@ -22,15 +22,15 @@ elts = (Float64, ComplexF64) Acopy = copy(A) Q, R = @constinferred qr_full(A, labels_A, labels_Q, labels_R) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) + A′ = contractalign(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) @test A ≈ A′ @test size(Q, 1) * size(Q, 2) == size(Q, 3) # Q is unitary Q, R = qr_full(A, (2, 1), (4, 3)) - @test A ≈ contract(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) + @test A ≈ contractalign(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) Q, R = qr_full(A, Val(2)) - @test A ≈ contract((:a, :b, :c, :d), Q, (:a, :b, :q), R, (:q, :c, :d)) + @test A ≈ contractalign((:a, :b, :c, :d), Q, (:a, :b, :q), R, (:q, :c, :d)) end @testset "Compact QR ($T)" for T in elts @@ -42,7 +42,7 @@ end Acopy = copy(A) Q, R = @constinferred qr_compact(A, labels_A, labels_Q, labels_R) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) + A′ = contractalign(labels_A, Q, (labels_Q..., :q), R, (:q, labels_R...)) @test A ≈ A′ @test size(Q, 3) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) end @@ -58,12 +58,12 @@ end Acopy = copy(A) L, Q = @constinferred lq_full(A, labels_A, labels_L, labels_Q) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) + A′ = contractalign(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) @test A ≈ A′ @test size(Q, 1) == size(Q, 2) * size(Q, 3) # Q is unitary L, Q = lq_full(A, (2, 1), (4, 3)) - @test A ≈ contract(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) + @test A ≈ contractalign(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) end @testset "Compact LQ ($T)" for T in elts @@ -75,7 +75,7 @@ end Acopy = copy(A) L, Q = @constinferred lq_compact(A, labels_A, labels_L, labels_Q) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) + A′ = contractalign(labels_A, L, (labels_L..., :q), Q, (:q, labels_Q...)) @test A ≈ A′ @test size(Q, 1) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) # Q is unitary end @@ -95,8 +95,8 @@ end # `D` is returned bare (the spectrum over the internal bond), which is a `Diagonal`. @test D isa Diagonal - AV = contract((:a, :b, :D), A, labels_A, V, (labels_V′..., :D)) - VD = contract((:a, :b, :D), V, (labels_V..., :D′), D, (:D′, :D)) + AV = contractalign((:a, :b, :D), A, labels_A, V, (labels_V′..., :D)) + VD = contractalign((:a, :b, :D), V, (labels_V..., :D′), D, (:D′, :D)) @test AV ≈ VD Dvals = eig_vals(A, labels_A, labels_V, labels_V′) @@ -118,8 +118,8 @@ end @test eltype(V) == eltype(A) @test D isa Diagonal - AV = contract((:a, :b, :D), A, labels_A, V, (labels_V′..., :D)) - VD = contract((:a, :b, :D), V, (labels_V..., :D′), D, (:D′, :D)) + AV = contractalign((:a, :b, :D), A, labels_A, V, (labels_V′..., :D)) + VD = contractalign((:a, :b, :D), V, (labels_V..., :D′), D, (:D′, :D)) @test AV ≈ VD Dvals = eigh_vals(A, labels_A, labels_V, labels_V′) @@ -139,26 +139,26 @@ end U, S, Vᴴ = @constinferred svd_full(A, labels_A, labels_U, labels_Vᴴ) @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (labels_U..., :u), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) @test A ≈ A′ @test size(U, 1) * size(U, 2) == size(U, 3) # U is unitary @test size(Vᴴ, 1) == size(Vᴴ, 2) * size(Vᴴ, 3) # V is unitary U, S, Vᴴ = svd_full(A, (2, 1), (4, 3)) US, labels_US = contract(U, (labels_U..., :u), S, (:u, :v)) - @test A ≈ contract(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) + @test A ≈ contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) U, S, Vᴴ = @constinferred svd_full(A, labels_A, labels_A, ()) @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (labels_A..., :u), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v,)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v,)) @test A ≈ A′ @test size(Vᴴ, 1) == 1 U, S, Vᴴ = @constinferred svd_full(A, labels_A, (), labels_A) @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (:u,), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v, labels_A...)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_A...)) @test A ≈ A′ @test size(U, 2) == 1 end @@ -174,7 +174,7 @@ end @test A == Acopy # should not have altered initial array @test S isa Diagonal US, labels_US = contract(U, (labels_U..., :u), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) @test A ≈ A′ k = min(size(S)...) @test size(U, 3) == k == size(Vᴴ, 1) @@ -185,14 +185,14 @@ end U, S, Vᴴ = @constinferred svd_compact(A, labels_A, labels_A, ()) @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (labels_A..., :u), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v,)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v,)) @test A ≈ A′ @test size(U, ndims(U)) == 1 == size(Vᴴ, 1) U, S, Vᴴ = @constinferred svd_compact(A, labels_A, (), labels_A) @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (:u,), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v, labels_A...)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_A...)) @test A ≈ A′ @test size(U, 1) == 1 == size(Vᴴ, 1) end @@ -212,7 +212,7 @@ end @test A == Acopy # should not have altered initial array US, labels_US = contract(U, (labels_U..., :u), S, (:u, :v)) - A′ = contract(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) + A′ = contractalign(labels_A, US, labels_US, Vᴴ, (:v, labels_Vᴴ...)) @test norm(A - A′) ≈ S_untrunc[end] @test size(S, 1) == size(S_untrunc, 1) - 1 # `ϵ` is the 2-norm of the discarded singular values (here the single dropped value). @@ -229,18 +229,36 @@ end N = @constinferred left_null(A, labels_A, labels_codomain, labels_domain) @test A == Acopy # should not have altered initial array # N^ba_n' * A^ba_dc = 0 - NA = contract((:n, labels_domain...), conj(N), (labels_codomain..., :n), A, labels_A) + NA = contractalign( + (:n, labels_domain...), + conj(N), + (labels_codomain..., :n), + A, + labels_A + ) @test norm(NA) ≈ 0 atol = 1.0e-14 NN = - contract((:n, :n′), conj(N), (labels_codomain..., :n), N, (labels_codomain..., :n′)) + contractalign( + (:n, :n′), + conj(N), + (labels_codomain..., :n), + N, + (labels_codomain..., :n′) + ) @test NN ≈ LinearAlgebra.I Nᴴ = @constinferred right_null(A, labels_A, labels_codomain, labels_domain) @test A == Acopy # should not have altered initial array # A^ba_dc * N^dc_n' = 0 - AN = contract((labels_codomain..., :n), A, labels_A, conj(Nᴴ), (:n, labels_domain...)) + AN = contractalign( + (labels_codomain..., :n), + A, + labels_A, + conj(Nᴴ), + (:n, labels_domain...) + ) @test norm(AN) ≈ 0 atol = 1.0e-14 - NN = contract((:n, :n′), Nᴴ, (:n, labels_domain...), Nᴴ, (:n′, labels_domain...)) + NN = contractalign((:n, :n′), Nᴴ, (:n, labels_domain...), Nᴴ, (:n′, labels_domain...)) end @testset "Left polar ($T)" for T in elts @@ -252,7 +270,7 @@ end Acopy = copy(A) W, P = left_polar(A, labels_A, labels_W, labels_P) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) + A′ = contractalign(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) @test A ≈ A′ @test size(W, 3) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) end @@ -266,7 +284,7 @@ end Acopy = copy(A) P, W = right_polar(A, labels_A, labels_P, labels_W) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) + A′ = contractalign(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) @test A ≈ A′ @test size(W, 1) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) end @@ -280,12 +298,12 @@ end Acopy = copy(A) W, P = left_orth(A, labels_A, labels_W, labels_P) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) + A′ = contractalign(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) @test A ≈ A′ @test size(W, 3) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) W, P = left_orth(A, (2, 1), (4, 3)) - @test A ≈ contract(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) + @test A ≈ contractalign(labels_A, W, (labels_W..., :w), P, (:w, labels_P...)) end @testset "Right orth ($T)" for T in elts @@ -297,12 +315,12 @@ end Acopy = copy(A) P, W = right_orth(A, labels_A, labels_P, labels_W) @test A == Acopy # should not have altered initial array - A′ = contract(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) + A′ = contractalign(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) @test A ≈ A′ @test size(W, 1) == min(size(A, 1) * size(A, 2), size(A, 3) * size(A, 4)) P, W = right_orth(A, (2, 1), (4, 3)) - @test A ≈ contract(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) + @test A ≈ contractalign(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) end # one (identity tensor) diff --git a/test/test_matricize.jl b/test/test_matricize.jl index 0582b4bd..1b2dd58c 100644 --- a/test/test_matricize.jl +++ b/test/test_matricize.jl @@ -108,7 +108,7 @@ end # for the identity destination bipermutation and the gather/scatter route otherwise. a1 = randn(rng, 2, 3, 5) a2 = randn(rng, 5, 3, 2) - ref = TensorAlgebra.contract((:i, :j, :k, :l), a1, (:i, :j, :m), a2, (:m, :k, :l)) + ref = TensorAlgebra.contractalign((:i, :j, :k, :l), a1, (:i, :j, :m), a2, (:m, :k, :l)) for labels in ((:i, :j, :k, :l), (:k, :l, :i, :j), (:k, :i, :l, :j)) perm = map(l -> findfirst(==(l), (:i, :j, :k, :l)), labels) dest = randn(rng, map(d -> size(ref, d), perm)...) diff --git a/test/test_mooncakeext.jl b/test/test_mooncakeext.jl index 840b5f3c..999f028e 100644 --- a/test/test_mooncakeext.jl +++ b/test/test_mooncakeext.jl @@ -2,7 +2,7 @@ using Mooncake: Mooncake using Random: Random using TensorAlgebra: BiTuple, ContractAlgorithm, DefaultContractAlgorithm, Matricize, allocate_output, biperm, biperms, check_input, contract, contract!, contract_labels, - contractadd!, default_contract_algorithm, select_contract_algorithm + contractadd!, contractpermadd!, default_contract_algorithm, select_contract_algorithm using Test: @test, @testset @testset "MooncakeExt" begin @@ -65,7 +65,7 @@ using Test: @test, @testset @testset "contract" begin α = true β = false - @testset "contractadd! (BiTuple)" begin + @testset "contractpermadd! (BiTuple)" begin dest = randn(elt, (2, 2)) a1 = randn(elt, (2, 2)) a2 = randn(elt, (2, 2)) @@ -73,7 +73,7 @@ using Test: @test, @testset biperm1 = BiTuple((1,), (2,)) biperm2 = BiTuple((1,), (2,)) Mooncake.TestUtils.test_rule( - rng, contractadd!, dest, biperm_dest.t1, biperm_dest.t2, + rng, contractpermadd!, dest, biperm_dest.t1, biperm_dest.t2, a1, biperm1.t1, biperm1.t2, a2, biperm2.t1, biperm2.t2, α, β; atol, rtol, mode, is_primitive ) From 678b3ba2e37e3de0b171185c1a27c2bab6e4c116 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 16 Sep 2026 01:35:06 -0400 Subject: [PATCH 16/33] Declare contractadd! public It was the one rung of the labels ladder that was neither exported nor public, though it is as much a user-facing entry point as the rest. Co-Authored-By: Claude Opus 5 (1M context) --- src/TensorAlgebra.jl | 2 +- test/test_exports.jl | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index b3b04fd7..d1248a09 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -10,7 +10,7 @@ export contract, contract!, contractalign, dual, eig_full, eig_trunc, eig_vals, if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public allocate_output, biperm, bipartition, cat_similar, check_input, concatenate, concatenate!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, output_axes, select_algorithm, default_algorithm, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" + "public allocate_output, biperm, bipartition, cat_similar, check_input, concatenate, concatenate!, ContractAlgorithm, contractadd!, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, output_axes, select_algorithm, default_algorithm, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" ) ) end diff --git a/test/test_exports.jl b/test/test_exports.jl index 89f67df8..40e6848f 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -39,7 +39,8 @@ using Test: @test, @testset exports, [ :allocate_output, :biperm, :bipartition, :cat_similar, :check_input, - :concatenate, :concatenate!, :ContractAlgorithm, :contractopadd!, + :concatenate, :concatenate!, :ContractAlgorithm, :contractadd!, + :contractopadd!, :contractperm, :contractperm!, :contractpermadd!, :contractpermopadd!, :data, :datatype, :directsum, From 6100505d196d547d72def8fc57082319ab8cf8c5 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 16 Sep 2026 09:17:57 -0400 Subject: [PATCH 17/33] Drop the algorithm-positional non-mutating contractions A contraction algorithm does not choose how the output is allocated, so these two entry points were a second way to contract that bypassed `allocate_output`. The algorithm stays a keyword above the in-place primitive. Co-Authored-By: Claude Opus 5 (1M context) --- .../TensorAlgebraTensorOperationsExt.jl | 35 ------------------- 1 file changed, 35 deletions(-) diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 79c84f3c..6017a1c7 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -21,41 +21,6 @@ end # Using TensorOperations backends as TensorAlgebra implementations # ---------------------------------------------------------------- -# not in-place -function TA.contractperm( - algorithm::TensorOperationsAlgorithm, - perm_dest_codomain, perm_dest_domain, - a1::AbstractArray, perm1_codomain, perm1_domain, - a2::AbstractArray, perm2_codomain, perm2_domain - ) - permblocks1 = Tuple.((perm1_codomain, perm1_domain)) - permblocks2 = Tuple.((perm2_codomain, perm2_domain)) - permblocks_dest = Tuple.((perm_dest_codomain, perm_dest_domain)) - conj1, conj2 = false, false - α = true - return TO.tensorcontract( - a1, permblocks1, conj1, a2, permblocks2, conj2, - permblocks_dest, α, backend(algorithm), allocator(algorithm) - ) -end - -function TA.contractalign( - algorithm::TensorOperationsAlgorithm, - labels_dest, - a1::AbstractArray, labels1, - a2::AbstractArray, labels2 - ) - permblocks1, permblocks2, permblocks_dest = - TO.contract_indices(labels1, labels2, labels_dest) - conj1, conj2 = false, false - α = true - return TO.tensorcontract( - a1, permblocks1, conj1, a2, permblocks2, conj2, - permblocks_dest, α, backend(algorithm), allocator(algorithm) - ) -end - -# in-place function TA.contractpermopadd!( algorithm::TensorOperationsAlgorithm, a_dest, perm_dest_codomain, perm_dest_domain, From ca7f06f22dabfb64d197bd6a4ad80a8cffdb972f Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Tue, 22 Sep 2026 20:20:26 -0400 Subject: [PATCH 18/33] Restore sqrth_invsqrth_safe Belief-propagation simple update needs both square roots of the same Hermitian matrix, and taking them from one eigendecomposition rather than two is worth the extra name. Co-Authored-By: Claude Opus 5 (1M context) --- src/MatrixAlgebra.jl | 29 +++++++++++++++++++++++++++++ src/TensorAlgebra.jl | 2 +- src/factorizations.jl | 37 +++++++++++++++++++++++++++++++++---- test/test_exports.jl | 2 ++ test/test_matrixalgebra.jl | 10 ++++++++++ 5 files changed, 75 insertions(+), 5 deletions(-) diff --git a/src/MatrixAlgebra.jl b/src/MatrixAlgebra.jl index 05e5aba6..f519acc0 100644 --- a/src/MatrixAlgebra.jl +++ b/src/MatrixAlgebra.jl @@ -6,6 +6,7 @@ export invsqrt_diag_safe, pow_diag_safe!, powh_safe, sqrt_diag_safe, + sqrth_invsqrth_safe, sqrth_safe using LinearAlgebra: LinearAlgebra, Diagonal, isdiag, norm @@ -162,6 +163,34 @@ $(_clamp_kwargs_doc("M")) """ invsqrth_safe(M; kwargs...) = powh_safe(M, -1 // 2; kwargs...) +""" + sqrth_invsqrth_safe(M; alg=nothing, atol=0, rtol=eps(real(eltype(M)))^(3//4)) -> M^(1//2), M^(-1//2) + +Square root and pseudo-inverse square root of a Hermitian positive +semi-definite matrix, from a single eigendecomposition. Equivalent +to `(sqrth_safe(M; ...), invsqrth_safe(M; ...))` but with the +eigendecomposition computed once. Eigenvalues below tolerance are clamped +to zero in both factors (Moore-Penrose convention for the inverse). + +The input must be Hermitian (as for `MatrixAlgebraKit.eigh_full`): project +with `MatrixAlgebraKit.project_hermitian` first if it is Hermitian only up +to numerical noise. + +## Keyword arguments + + - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. + +$(_clamp_kwargs_doc("M")) +""" +function sqrth_invsqrth_safe(M; alg = nothing, kwargs...) + if isdiag(M) + return pow_diag_safe(M, 1 // 2; kwargs...), pow_diag_safe(M, -1 // 2; kwargs...) + end + D, V = MAK.eigh_full(M; alg) + return V * pow_diag_safe(D, 1 // 2; kwargs...) * V', + V * pow_diag_safe(D, -1 // 2; kwargs...) * V' +end + using MatrixAlgebraKit: MatrixAlgebraKit, TruncationStrategy struct TruncationDegenerate{Strategy <: TruncationStrategy, T <: Real} <: TruncationStrategy diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index d1248a09..4c74aa0f 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -4,7 +4,7 @@ export contract, contract!, contractalign, dual, eig_full, eig_trunc, eig_vals, eigh_trunc, eigh_vals, invsqrth_safe, isdual, left_null, left_orth, left_polar, lq_compact, lq_full, project_hermitian, qr_compact, - qr_full, right_null, right_orth, right_polar, sqrth_safe, + qr_full, right_null, right_orth, right_polar, sqrth_invsqrth_safe, sqrth_safe, svd_compact, svd_full, svd_trunc, svd_vals if VERSION >= v"1.11.0-DEV.469" diff --git a/src/factorizations.jl b/src/factorizations.jl index 63440716..95d2fce2 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -63,7 +63,7 @@ end # Read-only tier: the matrix-level entries never mutate their input (they copy internally), so # the perm form consumes the maybe-alias `matricize` matricization directly. for f in ( - :sqrth_safe, :invsqrth_safe, + :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, ) @eval begin function $f( @@ -90,7 +90,7 @@ for f in ( :svd_compact, :svd_full, :svd_trunc, :svd_vals, :eigh_full, :eig_full, :eigh_trunc, :eig_trunc, :eigh_vals, :eig_vals, :left_null, :right_null, - :sqrth_safe, :invsqrth_safe, :project_hermitian, + :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, :project_hermitian, ) @eval begin function $f(style::MatricizeStyle, A, ndims_codomain::Val{K}; kwargs...) where {K} @@ -629,7 +629,7 @@ up to numerical noise. $(MatrixAlgebra._clamp_kwargs_doc("A")) -See also [`invsqrth_safe`](@ref) and +See also [`invsqrth_safe`](@ref), [`sqrth_invsqrth_safe`](@ref), and [`MatrixAlgebra.sqrth_safe`](@ref). """ sqrth_safe @@ -652,7 +652,7 @@ first if it is Hermitian only up to numerical noise. $(MatrixAlgebra._clamp_kwargs_doc("A")) -See also [`sqrth_safe`](@ref) and +See also [`sqrth_safe`](@ref), [`sqrth_invsqrth_safe`](@ref), and [`MatrixAlgebra.invsqrth_safe`](@ref). """ invsqrth_safe @@ -688,6 +688,35 @@ function unmatricize_factors( return unmatricize(style, H_mat, axes_codomain, axes_domain) end +""" + sqrth_invsqrth_safe(A, labels_A, labels_codomain, labels_domain; kwargs...) -> P, Pinv + sqrth_invsqrth_safe(A, perm_codomain, perm_domain; kwargs...) -> P, Pinv + sqrth_invsqrth_safe(A, ndims_codomain::Val; kwargs...) -> P, Pinv + +Square root and pseudo-inverse square root of a generic N-dimensional +array (see [`sqrth_safe`](@ref) and [`invsqrth_safe`](@ref)), from a +single eigendecomposition. Both results carry the same codomain and +domain axes as `A`. + +## Keyword arguments + + - `alg`: forwarded to `MatrixAlgebraKit.eigh_full`. + +$(MatrixAlgebra._clamp_kwargs_doc("A")) + +See also [`MatrixAlgebra.sqrth_invsqrth_safe`](@ref). +""" +sqrth_invsqrth_safe + +function unmatricize_factors( + ::typeof(sqrth_invsqrth_safe), style::MatricizeStyle, F, + axes_codomain, axes_domain + ) + P_mat, Pinv_mat = F + return unmatricize(style, P_mat, axes_codomain, axes_domain), + unmatricize(style, Pinv_mat, axes_codomain, axes_domain) +end + """ TensorAlgebra.one(A, labels_A, labels_codomain, labels_domain) -> Id TensorAlgebra.one(A, perm_codomain, perm_domain) -> Id diff --git a/test/test_exports.jl b/test/test_exports.jl index 40e6848f..dabdbb81 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -27,6 +27,7 @@ using Test: @test, @testset :right_null, :right_orth, :right_polar, + :sqrth_invsqrth_safe, :sqrth_safe, :svd_compact, :svd_full, @@ -64,6 +65,7 @@ using Test: @test, @testset :pow_diag_safe!, :powh_safe, :sqrt_diag_safe, + :sqrth_invsqrth_safe, :sqrth_safe, ] @test issetequal(names(TensorAlgebra.MatrixAlgebra), exports) diff --git a/test/test_matrixalgebra.jl b/test/test_matrixalgebra.jl index bda1a277..9eaea1f7 100644 --- a/test/test_matrixalgebra.jl +++ b/test/test_matrixalgebra.jl @@ -159,6 +159,16 @@ elts = (Float32, Float64, ComplexF32, ComplexF64) invsqrtA = MatrixAlgebra.invsqrth_safe(A) @test invsqrtA * sqrtA ≈ I(n) + # The paired form shares one eigendecomposition, so it has to agree with the + # separate calls, on both the dense and the `isdiag` fast paths. + P, Pinv = MatrixAlgebra.sqrth_invsqrth_safe(A) + @test P ≈ sqrtA + @test Pinv ≈ invsqrtA + Ddiag = Diagonal(rand(real(elt), n) .+ 1) + Pd, Pdinv = MatrixAlgebra.sqrth_invsqrth_safe(Matrix(Ddiag)) + @test Pd ≈ MatrixAlgebra.sqrth_safe(Matrix(Ddiag)) + @test Pdinv ≈ MatrixAlgebra.invsqrth_safe(Matrix(Ddiag)) + # Integer power: passes through without clamping affecting result. @test MatrixAlgebra.powh_safe(A, 2) ≈ A * A From 984618efa9e9323fb782d993e9f6249e4c05983e Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 08:10:40 -0400 Subject: [PATCH 19/33] Correct the contract and contractalign docstrings The `contract` example asserted a tuple of labels, but the surviving labels come back as a `Vector`, which is what makes the return type concrete. The `contractalign` description is cut to the contract it actually has. Co-Authored-By: Claude Opus 5 (1M context) --- src/contract/contract.jl | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/contract/contract.jl b/src/contract/contract.jl index a3d47e17..1c4edc99 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -20,7 +20,9 @@ julia> a, b = randn(2, 3), randn(3, 4); julia> ab, labels = contract(a, (:i, :j), b, (:j, :k)); julia> labels -(:i, :k) +2-element Vector{Symbol}: + :i + :k julia> ab ≈ a * b true @@ -45,11 +47,9 @@ end """ contractalign(labels_dest, a1, labels1, a2, labels2, ...; alg = nothing) -> a_dest -Contract the arrays over the labels they share into a result whose dimensions carry -`labels_dest`, which must be the surviving labels in some order. - -This is [`contract`](@ref) with the output specified, so it returns the array on its own. The -name matches `ITensorBase.align`: arrange the result's dimensions to match the labels given. +Contract the input arrays over the shared labels, aligning the output array according to the +specified destination labels `labels_dest`. `labels_dest` must match the uncontracted labels, +i.e. `issetequal(labels_dest, symdiff(labels1, labels2))` must be `true`. ```jldoctest julia> using TensorAlgebra: contractalign From 20d66f78b422c89500598cf31174249a191715e3 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 10:04:05 -0400 Subject: [PATCH 20/33] Declare the supported interface Loading TensorAlgebra alongside MatrixAlgebraKit made 21 bare factorization names ambiguous, so those move from `export` to `public`. The declaration now also covers the permute ladder and the rest of what GradedArrays and ITensorBase overload. Co-Authored-By: Claude Opus 5 (1M context) --- src/TensorAlgebra.jl | 9 ++--- test/test_exports.jl | 81 ++++++++++++++++++++++++-------------------- 2 files changed, 47 insertions(+), 43 deletions(-) diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 4c74aa0f..3c58556c 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -1,16 +1,11 @@ module TensorAlgebra -export contract, contract!, contractalign, dual, eig_full, eig_trunc, eig_vals, eigh_full, - eigh_trunc, - eigh_vals, invsqrth_safe, isdual, left_null, - left_orth, left_polar, lq_compact, lq_full, project_hermitian, qr_compact, - qr_full, right_null, right_orth, right_polar, sqrth_invsqrth_safe, sqrth_safe, - svd_compact, svd_full, svd_trunc, svd_vals +export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public allocate_output, biperm, bipartition, cat_similar, check_input, concatenate, concatenate!, ContractAlgorithm, contractadd!, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, directsum, flattenlinear, is_output_view, label_type, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, output_axes, select_algorithm, default_algorithm, permutedims, permutedims!, scalar, similar_map, TensorOperationsAlgorithm, to_range, tr, tryflattenlinear, ungrade, zero!, scale!, permuteddims, PermutedDims" + "public add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, default_algorithm, default_contract_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, Matricize, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsAlgorithm, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/test/test_exports.jl b/test/test_exports.jl index dabdbb81..3ddca1b6 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -1,3 +1,4 @@ +using MatrixAlgebraKit: MatrixAlgebraKit using TensorAlgebra: TensorAlgebra using Test: @test, @testset @@ -8,55 +9,63 @@ using Test: @test, @testset :contract!, :contractalign, :dual, - :eig_full, - :eig_trunc, - :eig_vals, - :eigh_full, - :eigh_trunc, - :eigh_vals, - :invsqrth_safe, :isdual, - :left_null, - :left_orth, - :left_polar, - :lq_compact, - :lq_full, - :project_hermitian, - :qr_compact, - :qr_full, - :right_null, - :right_orth, - :right_polar, - :sqrth_invsqrth_safe, - :sqrth_safe, - :svd_compact, - :svd_full, - :svd_trunc, - :svd_vals, + :MatrixAlgebra, ] # `public` (Julia 1.11+) adds names to `names()`; include them on 1.11+. if VERSION >= v"1.11.0-DEV.469" append!( exports, [ - :allocate_output, :biperm, :bipartition, :cat_similar, :check_input, - :concatenate, :concatenate!, :ContractAlgorithm, :contractadd!, - :contractopadd!, + :add!, :AddBroadcasted, :addends, :allocate_output, :allocate_project, + :arguments, :axes, :bipartition, :bipartition_axes, :biperm, + :bipermutedims, :bipermutedims!, :bipermutedimsopadd!, :cat_axis, + :cat_similar, :check_input, :concatenate, :concatenate!, + :ConjBroadcasted, :contractadd!, :ContractAlgorithm, :contractopadd!, :contractperm, :contractperm!, :contractpermadd!, :contractpermopadd!, - :data, - :datatype, :directsum, - :flattenlinear, :is_output_view, :label_type, - :matricize, :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, - :default_algorithm, :output_axes, :permutedims, :select_algorithm, - :permutedims!, :scalar, :similar_map, - :TensorOperationsAlgorithm, - :to_range, :tr, :tryflattenlinear, :ungrade, :zero!, :scale!, - :permuteddims, :PermutedDims, + :data, :datatype, :default_algorithm, :default_contract_algorithm, + :dims2cat, :directsum, :eig_full, :eig_trunc, :eig_vals, :eigh_full, + :eigh_trunc, :eigh_vals, :fill_map, :flattenlinear, :infer_aux_space, + :invsqrth_safe, :is_output_view, :is_projected, :isidentitybiperm, + :label_type, :left_null, :left_orth, :left_polar, :LinearBroadcasted, + :linearbroadcasted, :lq_compact, :lq_full, :Matricize, :matricize, + :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, + :MatricizeStyle, :MATRIX_FUNCTIONS, :ndims, :ndims_codomain, :one, + :ones_map, :operation, :output_axes, :PermutedDims, :permuteddims, + :permutedims, :permutedims!, :permutedimsadd!, :permutedimsop, + :permutedimsopadd!, :project, :project!, :project_aux, + :project_hermitian, :projectto!, :qr_compact, :qr_full, :rand_map, + :randn_map, :right_null, :right_orth, :right_polar, :scalar, :scale!, + :ScaledBroadcasted, :select_algorithm, :similar_map, :size, + :sqrth_invsqrth_safe, :sqrth_safe, :sum, :svd_compact, :svd_full, + :svd_trunc, :svd_vals, :TensorOperationsAlgorithm, :to_range, :tr, + :trivialrange, :tryflattenlinear, :tryproject, :tryproject_aux, + :unchecked_project, :unchecked_project_aux, :ungrade, :unmatricize, + :unmatricize!, :unmatricize_factors, :unproject, :unscaled, :zero!, + :zeros_map, ] ) end @test issetequal(names(TensorAlgebra), exports) + # The matrix-level factorizations are `public`, not exported: the names MatrixAlgebraKit + # also exports would otherwise collide in a session loading both packages, and the + # `MatrixAlgebra` spellings are reached through the submodule, which is exported. + exported = filter(n -> Base.isexported(TensorAlgebra, n), names(TensorAlgebra)) + @test issetequal( + exported, + [ + :TensorAlgebra, :contract, :contract!, :contractalign, :dual, :isdual, + :MatrixAlgebra, + ] + ) + @test isempty( + intersect( + exported, + filter(n -> Base.isexported(MatrixAlgebraKit, n), names(MatrixAlgebraKit)) + ) + ) + exports = [ :MatrixAlgebra, :invsqrt_diag_safe, From d06a5bdcc02923ca55b8c1f061212491fa45fe9b Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 10:07:14 -0400 Subject: [PATCH 21/33] Set the version to 0.21.0 Takes the prerelease suffix off so merging the release PR registers. Co-Authored-By: Claude Opus 5 (1M context) --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index 30d47b0b..3fe422fc 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.21.0-DEV" +version = "0.21.0" authors = ["ITensor developers and contributors"] [workspace] From 10be55e5728536b4d6b081c604b8ecb336519e6d Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 14:35:30 -0400 Subject: [PATCH 22/33] Generalize algorithm selection across operations Replaces the contract-specific algorithm selection with a generic layer keyed on the operation, so other operations can register defaults the same way. The algorithm types are renamed after the operation they implement. Co-Authored-By: Claude Opus 5 (1M context) --- .../TensorAlgebraMooncakeExt.jl | 10 ++-- ext/TensorAlgebraTensorKitExt.jl | 5 +- .../TensorAlgebraTensorOperationsExt.jl | 26 +++++---- src/TensorAlgebra.jl | 2 +- src/algorithm.jl | 46 +++++++-------- src/contract/contract.jl | 31 +++++----- src/contract/contract_matricize.jl | 2 +- src/contract/contractalgorithm.jl | 56 ++++--------------- test/test_basics.jl | 32 +++++------ test/test_exports.jl | 44 +++++++-------- test/test_matricizestyle.jl | 19 ++++--- test/test_mooncakeext.jl | 14 ++--- test/test_tensoroperations.jl | 26 ++++----- 13 files changed, 142 insertions(+), 171 deletions(-) diff --git a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl index 713a3e7d..89dda1d1 100644 --- a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl +++ b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl @@ -1,12 +1,12 @@ module TensorAlgebraMooncakeExt using Mooncake: Mooncake, @zero_derivative, DefaultCtx -using TensorAlgebra: BiTuple, ContractAlgorithm, allocate_output, biperm, biperms, +using TensorAlgebra: AbstractContractAlgorithm, BiTuple, allocate_output, biperm, biperms, check_input, contract, contract!, contract_labels, decode_contraction_labels, - default_contract_algorithm, encode_contraction_labels, select_contract_algorithm + default_algorithm, encode_contraction_labels, select_algorithm Mooncake.tangent_type(::Type{<:BiTuple}) = Mooncake.NoTangent -Mooncake.tangent_type(::Type{<:ContractAlgorithm}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:AbstractContractAlgorithm}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{ typeof(allocate_output), typeof(contract), Any, Any, Any, Any, Any, Any, Any, Any, @@ -32,7 +32,7 @@ Mooncake.tangent_type(::Type{<:ContractAlgorithm}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{typeof(contract_labels), Any, Any, Any, Any} @zero_derivative DefaultCtx Tuple{typeof(encode_contraction_labels), Any, Any} @zero_derivative DefaultCtx Tuple{typeof(decode_contraction_labels), Any, Any, Any} -@zero_derivative DefaultCtx Tuple{typeof(default_contract_algorithm), Any, Any} -@zero_derivative DefaultCtx Tuple{typeof(select_contract_algorithm), Any, Any, Any} +@zero_derivative DefaultCtx Tuple{typeof(default_algorithm), Any, Any, Any, Any} +@zero_derivative DefaultCtx Tuple{typeof(select_algorithm), Any, Any, Any, Any, Any} end diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 90bbb924..5f3e2262 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -307,10 +307,11 @@ function TensorAlgebra.add!( return VectorInterface.add!(y, x, α, β) end -function TensorAlgebra.default_contract_algorithm( +function TensorAlgebra.default_algorithm( + ::typeof(TensorAlgebra.contract!), ::Type{<:AbstractTensorMap}, ::Type{<:AbstractTensorMap}, ::Type{<:AbstractTensorMap} ) - return TensorAlgebra.ContractAlgorithm(TO.DefaultBackend()) + return TensorAlgebra.AbstractContractAlgorithm(TO.DefaultBackend()) end # ================================== linear-combination broadcast ========================= diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 6017a1c7..2a657647 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -1,28 +1,30 @@ module TensorAlgebraTensorOperationsExt -using TensorAlgebra: TensorAlgebra as TA, TensorOperationsAlgorithm +using TensorAlgebra: TensorAlgebra as TA, TensorOperationsContract using TensorOperations: TensorOperations as TO -# `TensorOperationsAlgorithm` stores `nothing` to mean "TensorOperations' default"; resolve +# `TensorOperationsContract` stores `nothing` to mean "TensorOperations' default"; resolve # those here, where the defaults can be named. -function backend(algorithm::TensorOperationsAlgorithm) +function backend(algorithm::TensorOperationsContract) return @something algorithm.backend TO.DefaultBackend() end -function allocator(algorithm::TensorOperationsAlgorithm) +function allocator(algorithm::TensorOperationsContract) return @something algorithm.allocator TO.DefaultAllocator() end -# Construct via the `ContractAlgorithm` public constructor seam as well. -TA.ContractAlgorithm(backend::TO.AbstractBackend) = TensorOperationsAlgorithm(; backend) -function TA.ContractAlgorithm(backend::TO.AbstractBackend, allocator) - return TensorOperationsAlgorithm(; backend, allocator) +# Construct via the `AbstractContractAlgorithm` public constructor seam as well. +function TA.AbstractContractAlgorithm(backend::TO.AbstractBackend) + return TensorOperationsContract(; backend) +end +function TA.AbstractContractAlgorithm(backend::TO.AbstractBackend, allocator) + return TensorOperationsContract(; backend, allocator) end # Using TensorOperations backends as TensorAlgebra implementations # ---------------------------------------------------------------- function TA.contractpermopadd!( - algorithm::TensorOperationsAlgorithm, + algorithm::TensorOperationsContract, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, @@ -49,7 +51,7 @@ function TO.tensorcontract!( a2::AbstractArray, permblocks2::TO.Index2Tuple, conj2::Bool, permblocks_dest::TO.Index2Tuple, α::Number, β::Number, - backend::TA.ContractAlgorithm, + backend::TA.AbstractContractAlgorithm, allocator ) op1 = conj1 ? conj : identity @@ -71,7 +73,7 @@ function TO.tensortrace!( permblocks_dest::TO.Index2Tuple, conj_src::Bool, α::Number, β::Number, - ::TA.ContractAlgorithm, + ::TA.AbstractContractAlgorithm, allocator ) return TO.tensortrace!( @@ -86,7 +88,7 @@ function TO.tensoradd!( permblocks_src::TO.Index2Tuple, conj_src::Bool, α::Number, β::Number, - ::TA.ContractAlgorithm, + ::TA.AbstractContractAlgorithm, allocator ) return TO.tensoradd!( diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 3c58556c..29b782b7 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -5,7 +5,7 @@ export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, default_algorithm, default_contract_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, Matricize, matricize, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsAlgorithm, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" + "public AbstractAlgorithm, AbstractContractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/src/algorithm.jl b/src/algorithm.jl index 28775bb8..a97a7556 100644 --- a/src/algorithm.jl +++ b/src/algorithm.jl @@ -1,16 +1,28 @@ """ - TensorAlgebra.select_algorithm(f, args...; alg = nothing, kwargs...) + TensorAlgebra.AbstractAlgorithm + +Supertype for the algorithm objects operations dispatch on. An operation's own supertype +subtypes this (for example `AbstractContractAlgorithm`), which is what makes an instance pass +through [`TensorAlgebra.select_algorithm`](@ref) unchanged. +""" +abstract type AbstractAlgorithm end + +""" + TensorAlgebra.select_algorithm(f, alg, args...) Resolve the algorithm operation `f` should run with on `args`. An `alg` of `nothing` defers to -[`TensorAlgebra.default_algorithm`](@ref); anything else is validated and passed through. +[`TensorAlgebra.default_algorithm`](@ref), an `AbstractAlgorithm` passes through unchanged, and +anything else is an error. -This is the forward-facing layer the per-operation resolvers sit under, so a caller that does not -care which operation it is dispatching writes `select_algorithm(f, ...)` and a backend registers -its choice on `default_algorithm(f, ...)`. +`alg` is positional so each operation can dispatch on the algorithm type. The user-facing entry +points take it as a keyword and hand it here. """ -function select_algorithm(f, args...; alg = nothing, kwargs...) - isnothing(alg) && return default_algorithm(f, args...; kwargs...) - return select_algorithm_specified(f, alg, args...; kwargs...) +select_algorithm(f, ::Nothing, args...) = default_algorithm(f, args...) +select_algorithm(f, alg::AbstractAlgorithm, args...) = alg +# `alg` named something that is not an algorithm at all. Reported against the operation rather +# than as a `MethodError` from inside a resolver. +function select_algorithm(f, alg, args...) + return throw(ArgumentError("`$alg` is not an algorithm for `$f`")) end """ @@ -20,20 +32,10 @@ end The algorithm operation `f` runs with on `args` when the caller names none. The types form is the registration point for a storage type; the values form defaults to it. -Each operation bridges to its own resolver, so `default_algorithm(contract, A1, A2)` is -[`TensorAlgebra.default_contract_algorithm`](@ref). +A storage type registers its choice per operation, so a backend that contracts its own way +adds a method to `default_algorithm(contract!, A_dest, A1, A2)`. """ -function default_algorithm(f, args...; kwargs...) - return default_algorithm(f, map(typeof, args)...; kwargs...) -end -function default_algorithm(f, argtypes::Type...; kwargs...) +default_algorithm(f, args...) = default_algorithm(f, map(typeof, args)...) +function default_algorithm(f, argtypes::Type...) return throw(MethodError(default_algorithm, (f, argtypes...))) end - -# `alg` named something. A resolved algorithm object passes through; anything else is a caller -# error, reported against the operation rather than as a `MethodError` from inside the resolver. -function select_algorithm_specified(f, alg, args...; kwargs...) - return throw( - ArgumentError("`$alg` is not an algorithm for `$f`") - ) -end diff --git a/src/contract/contract.jl b/src/contract/contract.jl index 1c4edc99..4e2b88e0 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -234,7 +234,7 @@ function contractpermopadd!( op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, α::Number, β::Number; - alg = nothing, kwargs... + alg = nothing ) check_input( contract!, @@ -242,7 +242,7 @@ function contractpermopadd!( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) - algorithm = select_algorithm(contract, a1, a2; alg, kwargs...) + algorithm = select_algorithm(contract!, alg, a_dest, a1, a2) return contractpermopadd!( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, @@ -254,7 +254,7 @@ end # contractpermopadd! (dispatched on the algorithm, bipartitioned permutations) # Required interface if not using matricized contraction function contractpermopadd!( - algorithm::ContractAlgorithm, + algorithm::AbstractContractAlgorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, @@ -274,17 +274,18 @@ function contractpermopadd!( ) end -# Bridges from the operation-generic algorithm layer in `algorithm.jl` down to the contraction -# resolvers. They live here rather than beside those resolvers because dispatching on -# `::typeof(contract)` needs `contract` to exist, and `contractalgorithm.jl` is included first for -# the algorithm types this file's signatures use. -function default_algorithm(::typeof(contract), A1::Type, A2::Type; kwargs...) - algorithm = default_contract_algorithm(A1, A2) - reject_algorithm_kwargs(algorithm; kwargs...) - return algorithm -end -function select_algorithm_specified( - ::typeof(contract), alg::ContractAlgorithm, a1, a2; kwargs... +# The contraction methods of the operation-generic algorithm layer in `algorithm.jl`. They live +# here rather than beside the algorithm types because dispatching on `::typeof(contract!)` needs +# it to exist, and `contractalgorithm.jl` is included first for the types these signatures use. +# Keyed on `contract!` rather than `contract` because the destination is in the signature, +# matching `check_input`. A backend can therefore choose on the destination, even though the +# generic default derives the matricization styles from the operands alone. +function default_algorithm( + ::typeof(contract!), A_dest::Type{<:AbstractArray}, + A1::Type{<:AbstractArray}, A2::Type{<:AbstractArray} ) - return select_contract_algorithm(alg, a1, a2; kwargs...) + return MatricizeContract(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) +end +function select_algorithm(::typeof(contract!), ::DefaultContractAlgorithm, a_dest, a1, a2) + return default_algorithm(contract!, a_dest, a1, a2) end diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 409d032d..7b4e66e3 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -1,7 +1,7 @@ using LinearAlgebra: mul! function contractpermopadd!( - algorithm::Matricize, + algorithm::MatricizeContract, a_dest::AbstractArray, biperm_dest_codomain, biperm_dest_domain, op1, a1::AbstractArray, biperm1_codomain, biperm1_domain, op2, a2::AbstractArray, biperm2_codomain, biperm2_domain, diff --git a/src/contract/contractalgorithm.jl b/src/contract/contractalgorithm.jl index c172ccff..924b142b 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -1,61 +1,27 @@ -abstract type ContractAlgorithm end -ContractAlgorithm(algorithm::ContractAlgorithm) = algorithm +abstract type AbstractContractAlgorithm <: AbstractAlgorithm end +AbstractContractAlgorithm(algorithm::AbstractContractAlgorithm) = algorithm -struct DefaultContractAlgorithm <: ContractAlgorithm end +struct DefaultContractAlgorithm <: AbstractContractAlgorithm end -struct Matricize{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm +struct MatricizeContract{LeftStyle, RightStyle, OutputStyle} <: AbstractContractAlgorithm left_matricize_style::LeftStyle right_matricize_style::RightStyle output_matricize_style::OutputStyle end -Matricize(matricize_style) = Matricize(matricize_style, matricize_style, matricize_style) -Matricize() = Matricize(ReshapeMatricize()) +function MatricizeContract(matricize_style) + return MatricizeContract(matricize_style, matricize_style, matricize_style) +end +MatricizeContract() = MatricizeContract(ReshapeMatricize()) """ - TensorOperationsAlgorithm(; backend = nothing, allocator = nothing) + TensorOperationsContract(; backend = nothing, allocator = nothing) Contract using TensorOperations, with `backend` selecting the contraction kernel and `allocator` the allocator for temporary tensors (e.g. `TensorOperations.ManualAllocator()`). A `nothing` field uses TensorOperations' default. Only usable with TensorOperations loaded. """ -Base.@kwdef struct TensorOperationsAlgorithm{Backend, Allocator} <: ContractAlgorithm +Base.@kwdef struct TensorOperationsContract{Backend, Allocator} <: + AbstractContractAlgorithm backend::Backend = nothing allocator::Allocator = nothing end - -# The contraction entry points collect trailing keywords and forward them here, so these accept -# `kwargs...` even though no `ContractAlgorithm` is configurable by keyword yet. Without it an -# unrecognized keyword surfaces as a `MethodError` on this internal function rather than as a -# complaint about the keyword the caller actually passed. -function reject_algorithm_kwargs(algorithm; kwargs...) - isempty(kwargs) && return nothing - names = join(map(k -> "`$k`", collect(keys(kwargs))), ", ") - return throw( - ArgumentError( - "unsupported keyword argument(s) $names for contraction algorithm `$(nameof(typeof(algorithm)))`" - ) - ) -end - -function select_contract_algorithm(algorithm, a1, a2; kwargs...) - return throw( - ArgumentError( - "`$algorithm` is not a contraction algorithm; pass a `ContractAlgorithm` as `alg`" - ) - ) -end -function select_contract_algorithm(algorithm::ContractAlgorithm, a1, a2; kwargs...) - reject_algorithm_kwargs(algorithm; kwargs...) - return algorithm -end -function select_contract_algorithm(algorithm::DefaultContractAlgorithm, a1, a2; kwargs...) - return default_contract_algorithm(a1, a2; kwargs...) -end -function default_contract_algorithm(a1, a2; kwargs...) - algorithm = default_contract_algorithm(typeof(a1), typeof(a2)) - reject_algorithm_kwargs(algorithm; kwargs...) - return algorithm -end -function default_contract_algorithm(A1::Type{<:AbstractArray}, A2::Type{<:AbstractArray}) - return Matricize(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) -end diff --git a/test/test_basics.jl b/test/test_basics.jl index 2ffabf85..72537b51 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,8 +1,8 @@ import TensorAlgebra using StableRNGs: StableRNG -using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, - contract!, contractadd!, contractalign, length_codomain, length_domain, matricize, - unmatricize, unmatricize! +using TensorAlgebra: AbstractContractAlgorithm, BiTuple, bipermutedims, bipermutedims!, + contract, contract!, contractadd!, contractalign, length_codomain, length_domain, + matricize, unmatricize, unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -170,25 +170,23 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @testset "contraction algorithm selection rejects unusable keywords" begin a1 = randn(2, 3) a2 = randn(3, 4) - # A keyword no algorithm can consume must name itself, not surface as a `MethodError` - # from inside the resolver. - @test_throws ArgumentError contractalign( - (1, 3), - a1, - (1, 2), - a2, - (2, 3); - nonsense = 1 + # Nothing in the algorithm-selection layer accepts keywords, so a keyword no + # algorithm can consume fails as a `MethodError`. + @test_throws MethodError contractalign( + (1, 3), a1, (1, 2), a2, (2, 3); nonsense = 1 ) - @test_throws ArgumentError contractalign( - (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize(), nonsense = 1 + @test_throws MethodError contractalign( + (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.MatricizeContract(), + nonsense = 1 ) # A non-algorithm passed as `alg` says so rather than erroring with "Not implemented". - @test_throws ArgumentError TensorAlgebra.select_contract_algorithm(:nope, a1, a2) + @test_throws ArgumentError TensorAlgebra.select_algorithm( + TensorAlgebra.contract!, :nope, a1 * a2, a1, a2 + ) # The supported spellings still work. @test contractalign((1, 3), a1, (1, 2), a2, (2, 3)) ≈ a1 * a2 @test contractalign( - (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.Matricize() + (1, 3), a1, (1, 2), a2, (2, 3); alg = TensorAlgebra.MatricizeContract() ) ≈ a1 * a2 end @@ -200,7 +198,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test a_dest == fill(2, (2, 2)) end - alg_tensoroperations = ContractAlgorithm(TensorOperations.StridedBLAS()) + alg_tensoroperations = AbstractContractAlgorithm(TensorOperations.StridedBLAS()) @testset "contract (eltype1=$elt1, eltype2=$elt2)" for elt1 in elts, elt2 in elts elt_dest = promote_type(elt1, elt2) a1 = ones(elt1, (1, 1)) diff --git a/test/test_exports.jl b/test/test_exports.jl index 3ddca1b6..caac6634 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -17,28 +17,28 @@ using Test: @test, @testset append!( exports, [ - :add!, :AddBroadcasted, :addends, :allocate_output, :allocate_project, - :arguments, :axes, :bipartition, :bipartition_axes, :biperm, - :bipermutedims, :bipermutedims!, :bipermutedimsopadd!, :cat_axis, - :cat_similar, :check_input, :concatenate, :concatenate!, - :ConjBroadcasted, :contractadd!, :ContractAlgorithm, :contractopadd!, - :contractperm, :contractperm!, :contractpermadd!, :contractpermopadd!, - :data, :datatype, :default_algorithm, :default_contract_algorithm, - :dims2cat, :directsum, :eig_full, :eig_trunc, :eig_vals, :eigh_full, - :eigh_trunc, :eigh_vals, :fill_map, :flattenlinear, :infer_aux_space, - :invsqrth_safe, :is_output_view, :is_projected, :isidentitybiperm, - :label_type, :left_null, :left_orth, :left_polar, :LinearBroadcasted, - :linearbroadcasted, :lq_compact, :lq_full, :Matricize, :matricize, - :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, - :MatricizeStyle, :MATRIX_FUNCTIONS, :ndims, :ndims_codomain, :one, - :ones_map, :operation, :output_axes, :PermutedDims, :permuteddims, - :permutedims, :permutedims!, :permutedimsadd!, :permutedimsop, - :permutedimsopadd!, :project, :project!, :project_aux, - :project_hermitian, :projectto!, :qr_compact, :qr_full, :rand_map, - :randn_map, :right_null, :right_orth, :right_polar, :scalar, :scale!, - :ScaledBroadcasted, :select_algorithm, :similar_map, :size, - :sqrth_invsqrth_safe, :sqrth_safe, :sum, :svd_compact, :svd_full, - :svd_trunc, :svd_vals, :TensorOperationsAlgorithm, :to_range, :tr, + :AbstractAlgorithm, :AbstractContractAlgorithm, :add!, :AddBroadcasted, + :addends, :allocate_output, :allocate_project, :arguments, :axes, + :bipartition, :bipartition_axes, :biperm, :bipermutedims, + :bipermutedims!, :bipermutedimsopadd!, :cat_axis, :cat_similar, + :check_input, :concatenate, :concatenate!, :ConjBroadcasted, + :contractadd!, :contractopadd!, :contractperm, :contractperm!, + :contractpermadd!, :contractpermopadd!, :data, :datatype, + :default_algorithm, :dims2cat, :directsum, :eig_full, :eig_trunc, + :eig_vals, :eigh_full, :eigh_trunc, :eigh_vals, :fill_map, + :flattenlinear, :infer_aux_space, :invsqrth_safe, :is_output_view, + :is_projected, :isidentitybiperm, :label_type, :left_null, :left_orth, + :left_polar, :LinearBroadcasted, :linearbroadcasted, :lq_compact, + :lq_full, :matricize, :MatricizeContract, :matricizeop, :matricizeop!, + :matricizeopcopy, :matricizeopview, :MatricizeStyle, :MATRIX_FUNCTIONS, + :ndims, :ndims_codomain, :one, :ones_map, :operation, :output_axes, + :PermutedDims, :permuteddims, :permutedims, :permutedims!, + :permutedimsadd!, :permutedimsop, :permutedimsopadd!, :project, + :project!, :project_aux, :project_hermitian, :projectto!, :qr_compact, + :qr_full, :rand_map, :randn_map, :right_null, :right_orth, :right_polar, + :scalar, :scale!, :ScaledBroadcasted, :select_algorithm, :similar_map, + :size, :sqrth_invsqrth_safe, :sqrth_safe, :sum, :svd_compact, :svd_full, + :svd_trunc, :svd_vals, :TensorOperationsContract, :to_range, :tr, :trivialrange, :tryflattenlinear, :tryproject, :tryproject_aux, :unchecked_project, :unchecked_project_aux, :ungrade, :unmatricize, :unmatricize!, :unmatricize_factors, :unproject, :unscaled, :zero!, diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index 3c623e2d..6976dbbe 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -1,5 +1,6 @@ using LinearAlgebra: I -using TensorAlgebra: TensorAlgebra as TA, Matricize, MatricizeStyle, ReshapeMatricize +using TensorAlgebra: + TensorAlgebra as TA, MatricizeContract, MatricizeStyle, ReshapeMatricize using Test: @test, @testset module MatricizeStyleTestUtils @@ -54,14 +55,14 @@ using .MatricizeStyleTestUtils: MyArray, MyArrayMatricize @test MatricizeStyle(MyArrayMatricize(), MyArrayMatricize()) ≡ MyArrayMatricize() @test MatricizeStyle(MyArrayMatricize(), ReshapeMatricize()) ≡ ReshapeMatricize() @test MatricizeStyle(ReshapeMatricize(), MyArrayMatricize()) ≡ ReshapeMatricize() - @test TA.default_contract_algorithm(typeof(a1), typeof(a1)) ≡ - Matricize(ReshapeMatricize()) - @test TA.default_contract_algorithm(typeof(a1), typeof(a2)) ≡ - Matricize(ReshapeMatricize()) - @test TA.default_contract_algorithm(typeof(a2), typeof(a1)) ≡ - Matricize(ReshapeMatricize()) - @test TA.default_contract_algorithm(typeof(a2), typeof(a2)) ≡ - Matricize(MyArrayMatricize()) + @test TA.default_algorithm(TA.contract!, typeof(a1), typeof(a1), typeof(a1)) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, typeof(a1), typeof(a1), typeof(a2)) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, typeof(a2), typeof(a2), typeof(a1)) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, typeof(a2), typeof(a2), typeof(a2)) ≡ + MatricizeContract(MyArrayMatricize()) end @testset "style threads through the unfold" begin diff --git a/test/test_mooncakeext.jl b/test/test_mooncakeext.jl index 999f028e..2d930d26 100644 --- a/test/test_mooncakeext.jl +++ b/test/test_mooncakeext.jl @@ -1,8 +1,8 @@ using Mooncake: Mooncake using Random: Random -using TensorAlgebra: BiTuple, ContractAlgorithm, DefaultContractAlgorithm, Matricize, - allocate_output, biperm, biperms, check_input, contract, contract!, contract_labels, - contractadd!, contractpermadd!, default_contract_algorithm, select_contract_algorithm +using TensorAlgebra: AbstractContractAlgorithm, BiTuple, DefaultContractAlgorithm, + MatricizeContract, allocate_output, biperm, biperms, check_input, contract, contract!, + contract_labels, contractadd!, contractpermadd!, default_algorithm, select_algorithm using Test: @test, @testset @testset "MooncakeExt" begin @@ -14,9 +14,9 @@ using Test: @test, @testset rtol = eps(real(elt))^(3 / 4) @testset "zero derivatives" begin @test Mooncake.tangent_type(BiTuple) ≡ Mooncake.NoTangent - @test Mooncake.tangent_type(ContractAlgorithm) ≡ Mooncake.NoTangent + @test Mooncake.tangent_type(AbstractContractAlgorithm) ≡ Mooncake.NoTangent @test Mooncake.tangent_type(DefaultContractAlgorithm) ≡ Mooncake.NoTangent - @test Mooncake.tangent_type(Matricize) ≡ Mooncake.NoTangent + @test Mooncake.tangent_type(MatricizeContract) ≡ Mooncake.NoTangent dest = randn(elt, (2, 2)) a1 = randn(elt, (2, 2)) @@ -55,10 +55,10 @@ using Test: @test, @testset rng, contract_labels, a1, labels1, a2, labels2; mode, is_primitive ) Mooncake.TestUtils.test_rule( - rng, default_contract_algorithm, a1, a2; mode, is_primitive + rng, default_algorithm, contract!, dest, a1, a2; mode, is_primitive ) Mooncake.TestUtils.test_rule( - rng, select_contract_algorithm, DefaultContractAlgorithm(), a1, a2; + rng, select_algorithm, contract!, DefaultContractAlgorithm(), dest, a1, a2; mode, is_primitive ) end diff --git a/test/test_tensoroperations.jl b/test/test_tensoroperations.jl index 68f9c9ee..23e784ec 100644 --- a/test/test_tensoroperations.jl +++ b/test/test_tensoroperations.jl @@ -1,5 +1,5 @@ -using TensorAlgebra: - ContractAlgorithm, Matricize, TensorOperationsAlgorithm, contract, contract! +using TensorAlgebra: AbstractContractAlgorithm, MatricizeContract, TensorOperationsContract, + contract, contract! using TensorOperations: @tensor, DefaultAllocator, DefaultBackend, ManualAllocator, ncon, tensorcontract using Test: @inferred, @test, @testset @@ -20,7 +20,7 @@ using Test: @inferred, @test, @testset false, ((1, 5, 3, 2, 4), ()), 1.0, - Matricize() + MatricizeContract() ) @test C1 ≈ C2 end @@ -38,14 +38,14 @@ elts = (Float32, Float64, ComplexF32, ComplexF64) @tensor HrA12[a, s1, s2, c] := rhoL[a, a'] * A1[a', t1, b] * A2[b, t2, c'] * rhoR[c', c] * H[s1, s2, t1, t2] - @tensor backend = Matricize() HrA12′[a, s1, s2, c] := + @tensor backend = MatricizeContract() HrA12′[a, s1, s2, c] := rhoL[a, a'] * A1[a', t1, b] * A2[b, t2, c'] * rhoR[c', c] * H[s1, s2, t1, t2] @test HrA12 ≈ HrA12′ @test HrA12 ≈ ncon( [rhoL, H, A2, rhoR, A1], [[-1, 1], [-2, -3, 4, 5], [2, 5, 3], [3, -4], [1, 4, 2]]; - backend = Matricize() + backend = MatricizeContract() ) E = @tensor rhoL[a', a] * A1[a, s, b] * @@ -54,7 +54,7 @@ elts = (Float32, Float64, ComplexF32, ComplexF64) H[t, t', s, s'] * conj(A1[a', t, b']) * conj(A2[b', t', c']) - @test E ≈ @tensor backend = Matricize() rhoL[a', a] * + @test E ≈ @tensor backend = MatricizeContract() rhoL[a', a] * A1[a, s, b] * A2[b, s', c] * rhoR[c, c'] * @@ -121,26 +121,26 @@ end ) tensors = map(splat(randn), sizes) result1 = ncon(tensors, indices) - result2 = ncon(tensors, indices; backend = Matricize()) + result2 = ncon(tensors, indices; backend = MatricizeContract()) @test result1 ≈ result2 end end -@testset "TensorOperationsAlgorithm allocator ($T)" for T in elts +@testset "TensorOperationsContract allocator ($T)" for T in elts a1 = randn(T, 4, 5, 3) a2 = randn(T, 3, 6) labels1 = (:i, :j, :k) labels2 = (:k, :l) ref, ref_labels = contract(a1, labels1, a2, labels2) - @test TensorOperationsAlgorithm() isa ContractAlgorithm + @test TensorOperationsContract() isa AbstractContractAlgorithm @testset "allocator = $(nameof(typeof(alloc)))" for alloc in ( DefaultAllocator(), ManualAllocator(), ) - alg = TensorOperationsAlgorithm(; allocator = alloc) + alg = TensorOperationsContract(; allocator = alloc) c, labels = contract(a1, labels1, a2, labels2; alg) @test labels == ref_labels @test c ≈ ref @@ -150,10 +150,10 @@ end @test c_dest ≈ ref end - # The `ContractAlgorithm(backend, allocator)` constructor seam. - seam = ContractAlgorithm(DefaultBackend(), ManualAllocator()) + # The `AbstractContractAlgorithm(backend, allocator)` constructor seam. + seam = AbstractContractAlgorithm(DefaultBackend(), ManualAllocator()) @test contract(a1, labels1, a2, labels2; alg = seam)[1] ≈ ref # `nothing` fields fall back to the TensorOperations defaults. - @test contract(a1, labels1, a2, labels2; alg = TensorOperationsAlgorithm())[1] ≈ ref + @test contract(a1, labels1, a2, labels2; alg = TensorOperationsContract())[1] ≈ ref end From c5324191394c653f24d565ca591637fa1fd081e7 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 15:02:38 -0400 Subject: [PATCH 23/33] Pack algorithm selection arguments into a tuple The value and type domains of `select_algorithm` and `default_algorithm` overlapped when an argument was itself a type. Packing the selection-relevant arguments into a tuple keeps `(Float64, Int)` and `Tuple{Float64, Int}` distinct. Co-Authored-By: Claude Opus 5 (1M context) --- .../TensorAlgebraMooncakeExt.jl | 4 +-- ext/TensorAlgebraTensorKitExt.jl | 4 +-- src/algorithm.jl | 28 +++++++++++-------- src/contract/contract.jl | 11 ++++---- test/test_basics.jl | 2 +- test/test_matricizestyle.jl | 8 +++--- test/test_mooncakeext.jl | 4 +-- 7 files changed, 32 insertions(+), 29 deletions(-) diff --git a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl index 89dda1d1..55fb7cfb 100644 --- a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl +++ b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl @@ -32,7 +32,7 @@ Mooncake.tangent_type(::Type{<:AbstractContractAlgorithm}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{typeof(contract_labels), Any, Any, Any, Any} @zero_derivative DefaultCtx Tuple{typeof(encode_contraction_labels), Any, Any} @zero_derivative DefaultCtx Tuple{typeof(decode_contraction_labels), Any, Any, Any} -@zero_derivative DefaultCtx Tuple{typeof(default_algorithm), Any, Any, Any, Any} -@zero_derivative DefaultCtx Tuple{typeof(select_algorithm), Any, Any, Any, Any, Any} +@zero_derivative DefaultCtx Tuple{typeof(default_algorithm), Any, Any} +@zero_derivative DefaultCtx Tuple{typeof(select_algorithm), Any, Any, Any} end diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 5f3e2262..25a05ca4 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -308,8 +308,8 @@ function TensorAlgebra.add!( end function TensorAlgebra.default_algorithm( - ::typeof(TensorAlgebra.contract!), ::Type{<:AbstractTensorMap}, - ::Type{<:AbstractTensorMap}, ::Type{<:AbstractTensorMap} + ::typeof(TensorAlgebra.contract!), + ::Type{<:Tuple{AbstractTensorMap, AbstractTensorMap, AbstractTensorMap}} ) return TensorAlgebra.AbstractContractAlgorithm(TO.DefaultBackend()) end diff --git a/src/algorithm.jl b/src/algorithm.jl index a97a7556..44746c62 100644 --- a/src/algorithm.jl +++ b/src/algorithm.jl @@ -8,7 +8,7 @@ through [`TensorAlgebra.select_algorithm`](@ref) unchanged. abstract type AbstractAlgorithm end """ - TensorAlgebra.select_algorithm(f, alg, args...) + TensorAlgebra.select_algorithm(f, alg, args::Tuple) Resolve the algorithm operation `f` should run with on `args`. An `alg` of `nothing` defers to [`TensorAlgebra.default_algorithm`](@ref), an `AbstractAlgorithm` passes through unchanged, and @@ -16,26 +16,30 @@ anything else is an error. `alg` is positional so each operation can dispatch on the algorithm type. The user-facing entry points take it as a keyword and hand it here. + +The selection-relevant arguments are packed into a tuple rather than spliced, so that the value +and type domains stay disjoint. `(Float64, Int)` is a pair of arguments that happen to be types, +while `Tuple{Float64, Int}` names the types of a pair of arguments. """ -select_algorithm(f, ::Nothing, args...) = default_algorithm(f, args...) -select_algorithm(f, alg::AbstractAlgorithm, args...) = alg +select_algorithm(f, ::Nothing, args::Tuple) = default_algorithm(f, args) +select_algorithm(f, alg::AbstractAlgorithm, args::Tuple) = alg # `alg` named something that is not an algorithm at all. Reported against the operation rather # than as a `MethodError` from inside a resolver. -function select_algorithm(f, alg, args...) +function select_algorithm(f, alg, args::Tuple) return throw(ArgumentError("`$alg` is not an algorithm for `$f`")) end """ - TensorAlgebra.default_algorithm(f, args...) - TensorAlgebra.default_algorithm(f, argtypes::Type...) + TensorAlgebra.default_algorithm(f, args::Tuple) + TensorAlgebra.default_algorithm(f, Args::Type{<:Tuple}) The algorithm operation `f` runs with on `args` when the caller names none. The types form is the -registration point for a storage type; the values form defaults to it. +registration point for a storage type, and the values form defaults to it. -A storage type registers its choice per operation, so a backend that contracts its own way -adds a method to `default_algorithm(contract!, A_dest, A1, A2)`. +A storage type registers its choice per operation, so a backend that contracts its own way adds a +method to `default_algorithm(contract!, ::Type{<:Tuple{A_dest, A1, A2}})`. """ -default_algorithm(f, args...) = default_algorithm(f, map(typeof, args)...) -function default_algorithm(f, argtypes::Type...) - return throw(MethodError(default_algorithm, (f, argtypes...))) +default_algorithm(f, args::Tuple) = default_algorithm(f, typeof(args)) +function default_algorithm(f, Args::Type{<:Tuple}) + return throw(MethodError(default_algorithm, (f, Args))) end diff --git a/src/contract/contract.jl b/src/contract/contract.jl index 4e2b88e0..a84a80e3 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -242,7 +242,7 @@ function contractpermopadd!( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) - algorithm = select_algorithm(contract!, alg, a_dest, a1, a2) + algorithm = select_algorithm(contract!, alg, (a_dest, a1, a2)) return contractpermopadd!( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, @@ -281,11 +281,10 @@ end # matching `check_input`. A backend can therefore choose on the destination, even though the # generic default derives the matricization styles from the operands alone. function default_algorithm( - ::typeof(contract!), A_dest::Type{<:AbstractArray}, - A1::Type{<:AbstractArray}, A2::Type{<:AbstractArray} - ) + ::typeof(contract!), ::Type{Tuple{A_dest, A1, A2}} + ) where {A_dest <: AbstractArray, A1 <: AbstractArray, A2 <: AbstractArray} return MatricizeContract(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) end -function select_algorithm(::typeof(contract!), ::DefaultContractAlgorithm, a_dest, a1, a2) - return default_algorithm(contract!, a_dest, a1, a2) +function select_algorithm(::typeof(contract!), ::DefaultContractAlgorithm, args::Tuple) + return default_algorithm(contract!, args) end diff --git a/test/test_basics.jl b/test/test_basics.jl index 72537b51..c7cb02f2 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -181,7 +181,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int ) # A non-algorithm passed as `alg` says so rather than erroring with "Not implemented". @test_throws ArgumentError TensorAlgebra.select_algorithm( - TensorAlgebra.contract!, :nope, a1 * a2, a1, a2 + TensorAlgebra.contract!, :nope, (a1 * a2, a1, a2) ) # The supported spellings still work. @test contractalign((1, 3), a1, (1, 2), a2, (2, 3)) ≈ a1 * a2 diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl index 6976dbbe..5dd25b17 100644 --- a/test/test_matricizestyle.jl +++ b/test/test_matricizestyle.jl @@ -55,13 +55,13 @@ using .MatricizeStyleTestUtils: MyArray, MyArrayMatricize @test MatricizeStyle(MyArrayMatricize(), MyArrayMatricize()) ≡ MyArrayMatricize() @test MatricizeStyle(MyArrayMatricize(), ReshapeMatricize()) ≡ ReshapeMatricize() @test MatricizeStyle(ReshapeMatricize(), MyArrayMatricize()) ≡ ReshapeMatricize() - @test TA.default_algorithm(TA.contract!, typeof(a1), typeof(a1), typeof(a1)) ≡ + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a1), typeof(a1), typeof(a1)}) ≡ MatricizeContract(ReshapeMatricize()) - @test TA.default_algorithm(TA.contract!, typeof(a1), typeof(a1), typeof(a2)) ≡ + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a1), typeof(a1), typeof(a2)}) ≡ MatricizeContract(ReshapeMatricize()) - @test TA.default_algorithm(TA.contract!, typeof(a2), typeof(a2), typeof(a1)) ≡ + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a2), typeof(a2), typeof(a1)}) ≡ MatricizeContract(ReshapeMatricize()) - @test TA.default_algorithm(TA.contract!, typeof(a2), typeof(a2), typeof(a2)) ≡ + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a2), typeof(a2), typeof(a2)}) ≡ MatricizeContract(MyArrayMatricize()) end diff --git a/test/test_mooncakeext.jl b/test/test_mooncakeext.jl index 2d930d26..3de93e0d 100644 --- a/test/test_mooncakeext.jl +++ b/test/test_mooncakeext.jl @@ -55,10 +55,10 @@ using Test: @test, @testset rng, contract_labels, a1, labels1, a2, labels2; mode, is_primitive ) Mooncake.TestUtils.test_rule( - rng, default_algorithm, contract!, dest, a1, a2; mode, is_primitive + rng, default_algorithm, contract!, (dest, a1, a2); mode, is_primitive ) Mooncake.TestUtils.test_rule( - rng, select_algorithm, contract!, DefaultContractAlgorithm(), dest, a1, a2; + rng, select_algorithm, contract!, DefaultContractAlgorithm(), (dest, a1, a2); mode, is_primitive ) end From 4b6542fb6658e0ac5dbfd8c099bff7e2f3fb1be0 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 16:34:10 -0400 Subject: [PATCH 24/33] Build the identity bipermutation from ranks alone Takes the codomain rank and the total instead of an array, which was only ever read for its `ndims`, matching the shape `bipartition` already uses. The factorization and matrix-function forwarders open-coded the same tuples and now call it. Co-Authored-By: Claude Opus 5 (1M context) --- src/bituple.jl | 13 ++++++++----- src/factorizations.jl | 19 +++++++++---------- src/matricize.jl | 16 +++++----------- src/matrixfunctions.jl | 7 ++----- 4 files changed, 24 insertions(+), 31 deletions(-) diff --git a/src/bituple.jl b/src/bituple.jl index d9344a55..604cee00 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -84,9 +84,12 @@ function isidentitybiperm(perm_codomain, perm_domain) return perm == ntuple(identity, length(perm)) end -# The identity bipermutation for a rank-`N` array split after `ndims_codomain` dimensions, i.e. -# the one `isidentitybiperm` accepts. The split-only `Val` conveniences build it to reach the -# bipermutation forms. -function identitybiperm(a, ndims_codomain::Val{K}) where {K} - return ntuple(identity, Val(K)), ntuple(i -> K + i, Val(ndims(a) - K)) +# The identity bipermutation splitting `K` dimensions after `K_codomain`, i.e. the one +# `isidentitybiperm` accepts. The split-only `Val` conveniences build it to reach the +# bipermutation forms. Takes the codomain rank and the total, the same shape as `bipartition`, +# which reads the total off the container being split. Both are `Val`s so the tuple lengths stay +# compile-time constants and the result is inferrable. +function identitybiperm(::Val{K_codomain}, ::Val{K}) where {K_codomain, K} + return ntuple(identity, Val(K_codomain)), + ntuple(i -> K_codomain + i, Val(K - K_codomain)) end diff --git a/src/factorizations.jl b/src/factorizations.jl index 95d2fce2..7112d030 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -93,12 +93,9 @@ for f in ( :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, :project_hermitian, ) @eval begin - function $f(style::MatricizeStyle, A, ndims_codomain::Val{K}; kwargs...) where {K} + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) return $f( - style, A, - ntuple(identity, ndims_codomain), - ntuple(i -> K + i, Val(ndims(A) - K)); - kwargs... + style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...; kwargs... ) end function $f(A, ndims_codomain::Val; kwargs...) @@ -172,7 +169,9 @@ true ``` """ function tr(style::MatricizeStyle, A, ndims_codomain::Val) - return LinearAlgebra.tr(matricize(style, A, identitybiperm(A, ndims_codomain)...)) + return LinearAlgebra.tr( + matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) + ) end function tr(A, ndims_codomain::Val) return tr(MatricizeStyle(A), A, ndims_codomain) @@ -558,7 +557,7 @@ The output satisfies `N' * A ≈ 0` and `N' * N ≈ I`. left_null function left_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) + A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) N = MatrixAlgebraKit.left_null!(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return unmatricize(style, N, axes_codomain, (conj(axes(N, ndims(N))),)) @@ -595,7 +594,7 @@ The output satisfies `A * Nᴴ' ≈ 0` and `Nᴴ * Nᴴ' ≈ I`. right_null function right_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) + A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) Nᴴ = MatrixAlgebraKit.right_null!(A_mat; kwargs...) _, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, Nᴴ, (axes(Nᴴ, 1),), axes_domain) @@ -751,7 +750,7 @@ true function one end function one!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(A, ndims_codomain)...) + A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) MatrixAlgebraKit.one!(A_mat) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, A_mat, axes_codomain, axes_domain) @@ -764,7 +763,7 @@ end # matricization directly when the style declares one at this split, and otherwise fills a # gathered matrix and scatters it back with `unmatricize!`. function one!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - perm_codomain, perm_domain = identitybiperm(A, ndims_codomain) + perm_codomain, perm_domain = identitybiperm(ndims_codomain, Val(ndims(A))) if is_output_view(matricizeop, style, identity, A, perm_codomain, perm_domain) MatrixAlgebraKit.one!( matricizeopview(style, identity, A, perm_codomain, perm_domain) diff --git a/src/matricize.jl b/src/matricize.jl index 830872ee..757dad8f 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -130,16 +130,16 @@ end # the bipermutation forms, not a dispatch tier. A style implements the hooks above and never these, # which is what keeps the copy path from recursing back through the router. function matricize(a, ndims_codomain::Val) - return matricize(a, identitybiperm(a, ndims_codomain)...) + return matricize(a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end function matricize(style::MatricizeStyle, a, ndims_codomain::Val) - return matricize(style, a, identitybiperm(a, ndims_codomain)...) + return matricize(style, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end function matricizeop(op, a, ndims_codomain::Val) - return matricizeop(op, a, identitybiperm(a, ndims_codomain)...) + return matricizeop(op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end function matricizeop(style::MatricizeStyle, op, a, ndims_codomain::Val) - return matricizeop(style, op, a, identitybiperm(a, ndims_codomain)...) + return matricizeop(style, op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end # Total: always fresh storage the caller owns. @@ -244,14 +244,8 @@ end # split at `ndims_codomain`. The split applies no permutation, so this is the bipermutation form at the # trivial bipermutation, reusing its in-place block scatter (no intermediate `unmatricize` copy). function unmatricize!(style::MatricizeStyle, a_dest, m, ndims_codomain::Val) - K = unval(ndims_codomain) - N = ndims(a_dest) return unmatricize!( - style, - a_dest, - m, - ntuple(identity, Val(K)), - ntuple(i -> K + i, Val(N - K)) + style, a_dest, m, identitybiperm(ndims_codomain, Val(ndims(a_dest)))... ) end function unmatricize!(a_dest, m, ndims_codomain::Val) diff --git a/src/matrixfunctions.jl b/src/matrixfunctions.jl index 4b5f5244..a128f22f 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -36,12 +36,9 @@ const MATRIX_FUNCTIONS = [ # the eager `bipermutedims` copy at the identity bipermutation. for f in MATRIX_FUNCTIONS @eval begin - function $f(style::MatricizeStyle, a, ndims_codomain::Val{K}; kwargs...) where {K} + function $f(style::MatricizeStyle, a, ndims_codomain::Val; kwargs...) return $f( - style, a, - ntuple(identity, ndims_codomain), - ntuple(i -> K + i, Val(ndims(a) - K)); - kwargs... + style, a, identitybiperm(ndims_codomain, Val(ndims(a)))...; kwargs... ) end function $f(a, ndims_codomain::Val; kwargs...) From 0f6924baf8278dfdb57e79cbbe2cc31cf51fa192 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Wed, 23 Sep 2026 16:34:24 -0400 Subject: [PATCH 25/33] Name each contract argument form Gives the bipermutation tier the same argument-form naming as the labels tier, with `contractpermalign` for the destination-specifying form. Also drops the variadic forms over three or more operands and shortens `AbstractContractAlgorithm` to `ContractAlgorithm`. Co-Authored-By: Claude Opus 5 (1M context) --- .../TensorAlgebraMooncakeExt.jl | 4 +- ext/TensorAlgebraTensorKitExt.jl | 2 +- .../TensorAlgebraTensorOperationsExt.jl | 12 +++--- src/TensorAlgebra.jl | 2 +- src/algorithm.jl | 2 +- src/contract/contract.jl | 41 ++++--------------- src/contract/contractalgorithm.jl | 10 ++--- test/test_basics.jl | 8 ++-- test/test_exports.jl | 8 ++-- test/test_mooncakeext.jl | 4 +- test/test_tensoroperations.jl | 10 ++--- 11 files changed, 40 insertions(+), 63 deletions(-) diff --git a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl index 55fb7cfb..99961560 100644 --- a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl +++ b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl @@ -1,12 +1,12 @@ module TensorAlgebraMooncakeExt using Mooncake: Mooncake, @zero_derivative, DefaultCtx -using TensorAlgebra: AbstractContractAlgorithm, BiTuple, allocate_output, biperm, biperms, +using TensorAlgebra: BiTuple, ContractAlgorithm, allocate_output, biperm, biperms, check_input, contract, contract!, contract_labels, decode_contraction_labels, default_algorithm, encode_contraction_labels, select_algorithm Mooncake.tangent_type(::Type{<:BiTuple}) = Mooncake.NoTangent -Mooncake.tangent_type(::Type{<:AbstractContractAlgorithm}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:ContractAlgorithm}) = Mooncake.NoTangent @zero_derivative DefaultCtx Tuple{ typeof(allocate_output), typeof(contract), Any, Any, Any, Any, Any, Any, Any, Any, diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 25a05ca4..a9114ceb 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -311,7 +311,7 @@ function TensorAlgebra.default_algorithm( ::typeof(TensorAlgebra.contract!), ::Type{<:Tuple{AbstractTensorMap, AbstractTensorMap, AbstractTensorMap}} ) - return TensorAlgebra.AbstractContractAlgorithm(TO.DefaultBackend()) + return TensorAlgebra.ContractAlgorithm(TO.DefaultBackend()) end # ================================== linear-combination broadcast ========================= diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 2a657647..0bcb5134 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -12,11 +12,11 @@ function allocator(algorithm::TensorOperationsContract) return @something algorithm.allocator TO.DefaultAllocator() end -# Construct via the `AbstractContractAlgorithm` public constructor seam as well. -function TA.AbstractContractAlgorithm(backend::TO.AbstractBackend) +# Construct via the `ContractAlgorithm` public constructor seam as well. +function TA.ContractAlgorithm(backend::TO.AbstractBackend) return TensorOperationsContract(; backend) end -function TA.AbstractContractAlgorithm(backend::TO.AbstractBackend, allocator) +function TA.ContractAlgorithm(backend::TO.AbstractBackend, allocator) return TensorOperationsContract(; backend, allocator) end @@ -51,7 +51,7 @@ function TO.tensorcontract!( a2::AbstractArray, permblocks2::TO.Index2Tuple, conj2::Bool, permblocks_dest::TO.Index2Tuple, α::Number, β::Number, - backend::TA.AbstractContractAlgorithm, + backend::TA.ContractAlgorithm, allocator ) op1 = conj1 ? conj : identity @@ -73,7 +73,7 @@ function TO.tensortrace!( permblocks_dest::TO.Index2Tuple, conj_src::Bool, α::Number, β::Number, - ::TA.AbstractContractAlgorithm, + ::TA.ContractAlgorithm, allocator ) return TO.tensortrace!( @@ -88,7 +88,7 @@ function TO.tensoradd!( permblocks_src::TO.Index2Tuple, conj_src::Bool, α::Number, β::Number, - ::TA.AbstractContractAlgorithm, + ::TA.ContractAlgorithm, allocator ) return TO.tensoradd!( diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 29b782b7..d2ba8c7c 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -5,7 +5,7 @@ export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public AbstractAlgorithm, AbstractContractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" + "public AbstractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermalign, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/src/algorithm.jl b/src/algorithm.jl index 44746c62..7a505ad1 100644 --- a/src/algorithm.jl +++ b/src/algorithm.jl @@ -2,7 +2,7 @@ TensorAlgebra.AbstractAlgorithm Supertype for the algorithm objects operations dispatch on. An operation's own supertype -subtypes this (for example `AbstractContractAlgorithm`), which is what makes an instance pass +subtypes this (for example `ContractAlgorithm`), which is what makes an instance pass through [`TensorAlgebra.select_algorithm`](@ref) unchanged. """ abstract type AbstractAlgorithm end diff --git a/src/contract/contract.jl b/src/contract/contract.jl index a84a80e3..1cf720ed 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -3,15 +3,12 @@ # contract (labels) """ - contract(a1, labels1, a2, labels2, ...; alg = nothing) -> a_dest, labels_dest + contract(a1, labels1, a2, labels2; alg = nothing) -> a_dest, labels_dest Contract the arrays over the labels they share, returning the result along with the labels of its dimensions. A label appearing on two operands is summed over, one appearing on a single operand survives, and `labels_dest` reports the surviving labels in the order the result carries them. -Operands past the second are contracted one pair at a time from left to right, so the call -expresses the contraction order rather than requesting an optimized one. - ```jldoctest julia> using TensorAlgebra: contract @@ -38,14 +35,9 @@ function contract(a1, labels1, a2, labels2; kwargs...) a_dest = contractalign(l_dest, a1, l1, a2, l2; kwargs...) return a_dest, decode_contraction_labels(l_dest, labels1, labels2) end -function contract(a1, labels1, a2, labels2, a3, labels3, rest...; kwargs...) - check_alternating_labels(contract, rest) - a12, labels12 = contract(a1, labels1, a2, labels2; kwargs...) - return contract(a12, labels12, a3, labels3, rest...; kwargs...) -end """ - contractalign(labels_dest, a1, labels1, a2, labels2, ...; alg = nothing) -> a_dest + contractalign(labels_dest, a1, labels1, a2, labels2; alg = nothing) -> a_dest Contract the input arrays over the shared labels, aligning the output array according to the specified destination labels `labels_dest`. `labels_dest` must match the uncontracted labels, @@ -79,37 +71,20 @@ function contractalign( kwargs... ) end -# Only the last pair lands on the requested output labels; the ones before it infer their own. -function contractalign( - labels_dest, a1, labels1, a2, labels2, a3, labels3, rest...; kwargs... - ) - check_alternating_labels(contractalign, rest) - a12, labels12 = contract(a1, labels1, a2, labels2; kwargs...) - return contractalign(labels_dest, a12, labels12, a3, labels3, rest...; kwargs...) -end function _contractalign( ::Val{K}, labels_dest, a1, labels1, a2, labels2, contracted1; kwargs... ) where {K} biperm_dest, biperm1, biperm2 = biperms(contract, Val(K), labels_dest, labels1, labels2, contracted1) - return contractperm(biperm_dest..., a1, biperm1..., a2, biperm2...; kwargs...) -end - -# The variadic forms take arrays and labels in alternating positions, so a trailing group with an -# odd length is a miscount at the call site rather than something to diagnose further down. -function check_alternating_labels(f, rest::Tuple) - iseven(length(rest)) || throw( - ArgumentError( - "`$f` takes each array followed by its labels, so the trailing arguments must come in pairs" - ) + return contractpermalign( + biperm_dest..., a1, biperm1..., a2, biperm2...; kwargs... ) - return nothing end # contractperm (bipartitioned permutations) # `perm` marks the whole biperm ladder: every rung has a labels-form sibling under the plain name, -# and once `contract` is variadic over operands the two can no longer be told apart by arity. +# so the suffix says which form a call site is in without the reader counting arguments. function contractperm( a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; @@ -119,14 +94,14 @@ function contractperm( Ndest = Val(length(perm1_codomain) + length(perm2_domain)) perm_dest_codomain, perm_dest_domain = bipartition(ntuple(identity, Ndest), Ndest_codomain) - return contractperm( + return contractpermalign( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; kwargs... ) end -function contractperm( +function contractpermalign( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; @@ -254,7 +229,7 @@ end # contractpermopadd! (dispatched on the algorithm, bipartitioned permutations) # Required interface if not using matricized contraction function contractpermopadd!( - algorithm::AbstractContractAlgorithm, + algorithm::ContractAlgorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, diff --git a/src/contract/contractalgorithm.jl b/src/contract/contractalgorithm.jl index 924b142b..e7ffc8a7 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -1,9 +1,9 @@ -abstract type AbstractContractAlgorithm <: AbstractAlgorithm end -AbstractContractAlgorithm(algorithm::AbstractContractAlgorithm) = algorithm +abstract type ContractAlgorithm <: AbstractAlgorithm end +ContractAlgorithm(algorithm::ContractAlgorithm) = algorithm -struct DefaultContractAlgorithm <: AbstractContractAlgorithm end +struct DefaultContractAlgorithm <: ContractAlgorithm end -struct MatricizeContract{LeftStyle, RightStyle, OutputStyle} <: AbstractContractAlgorithm +struct MatricizeContract{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm left_matricize_style::LeftStyle right_matricize_style::RightStyle output_matricize_style::OutputStyle @@ -21,7 +21,7 @@ Contract using TensorOperations, with `backend` selecting the contraction kernel A `nothing` field uses TensorOperations' default. Only usable with TensorOperations loaded. """ Base.@kwdef struct TensorOperationsContract{Backend, Allocator} <: - AbstractContractAlgorithm + ContractAlgorithm backend::Backend = nothing allocator::Allocator = nothing end diff --git a/test/test_basics.jl b/test/test_basics.jl index c7cb02f2..1903aab2 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,8 +1,8 @@ import TensorAlgebra using StableRNGs: StableRNG -using TensorAlgebra: AbstractContractAlgorithm, BiTuple, bipermutedims, bipermutedims!, - contract, contract!, contractadd!, contractalign, length_codomain, length_domain, - matricize, unmatricize, unmatricize! +using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, + contract!, contractadd!, contractalign, length_codomain, length_domain, matricize, + unmatricize, unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -198,7 +198,7 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int @test a_dest == fill(2, (2, 2)) end - alg_tensoroperations = AbstractContractAlgorithm(TensorOperations.StridedBLAS()) + alg_tensoroperations = ContractAlgorithm(TensorOperations.StridedBLAS()) @testset "contract (eltype1=$elt1, eltype2=$elt2)" for elt1 in elts, elt2 in elts elt_dest = promote_type(elt1, elt2) a1 = ones(elt1, (1, 1)) diff --git a/test/test_exports.jl b/test/test_exports.jl index caac6634..e945e1c9 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -17,13 +17,15 @@ using Test: @test, @testset append!( exports, [ - :AbstractAlgorithm, :AbstractContractAlgorithm, :add!, :AddBroadcasted, + :AbstractAlgorithm, :add!, :AddBroadcasted, :addends, :allocate_output, :allocate_project, :arguments, :axes, :bipartition, :bipartition_axes, :biperm, :bipermutedims, :bipermutedims!, :bipermutedimsopadd!, :cat_axis, :cat_similar, :check_input, :concatenate, :concatenate!, :ConjBroadcasted, - :contractadd!, :contractopadd!, :contractperm, :contractperm!, - :contractpermadd!, :contractpermopadd!, :data, :datatype, + :contractadd!, :ContractAlgorithm, :contractopadd!, :contractperm, + :contractperm!, + :contractpermadd!, :contractpermalign, :contractpermopadd!, :data, + :datatype, :default_algorithm, :dims2cat, :directsum, :eig_full, :eig_trunc, :eig_vals, :eigh_full, :eigh_trunc, :eigh_vals, :fill_map, :flattenlinear, :infer_aux_space, :invsqrth_safe, :is_output_view, diff --git a/test/test_mooncakeext.jl b/test/test_mooncakeext.jl index 3de93e0d..443df37a 100644 --- a/test/test_mooncakeext.jl +++ b/test/test_mooncakeext.jl @@ -1,6 +1,6 @@ using Mooncake: Mooncake using Random: Random -using TensorAlgebra: AbstractContractAlgorithm, BiTuple, DefaultContractAlgorithm, +using TensorAlgebra: BiTuple, ContractAlgorithm, DefaultContractAlgorithm, MatricizeContract, allocate_output, biperm, biperms, check_input, contract, contract!, contract_labels, contractadd!, contractpermadd!, default_algorithm, select_algorithm using Test: @test, @testset @@ -14,7 +14,7 @@ using Test: @test, @testset rtol = eps(real(elt))^(3 / 4) @testset "zero derivatives" begin @test Mooncake.tangent_type(BiTuple) ≡ Mooncake.NoTangent - @test Mooncake.tangent_type(AbstractContractAlgorithm) ≡ Mooncake.NoTangent + @test Mooncake.tangent_type(ContractAlgorithm) ≡ Mooncake.NoTangent @test Mooncake.tangent_type(DefaultContractAlgorithm) ≡ Mooncake.NoTangent @test Mooncake.tangent_type(MatricizeContract) ≡ Mooncake.NoTangent diff --git a/test/test_tensoroperations.jl b/test/test_tensoroperations.jl index 23e784ec..3f025443 100644 --- a/test/test_tensoroperations.jl +++ b/test/test_tensoroperations.jl @@ -1,5 +1,5 @@ -using TensorAlgebra: AbstractContractAlgorithm, MatricizeContract, TensorOperationsContract, - contract, contract! +using TensorAlgebra: + ContractAlgorithm, MatricizeContract, TensorOperationsContract, contract, contract! using TensorOperations: @tensor, DefaultAllocator, DefaultBackend, ManualAllocator, ncon, tensorcontract using Test: @inferred, @test, @testset @@ -133,7 +133,7 @@ end labels2 = (:k, :l) ref, ref_labels = contract(a1, labels1, a2, labels2) - @test TensorOperationsContract() isa AbstractContractAlgorithm + @test TensorOperationsContract() isa ContractAlgorithm @testset "allocator = $(nameof(typeof(alloc)))" for alloc in ( @@ -150,8 +150,8 @@ end @test c_dest ≈ ref end - # The `AbstractContractAlgorithm(backend, allocator)` constructor seam. - seam = AbstractContractAlgorithm(DefaultBackend(), ManualAllocator()) + # The `ContractAlgorithm(backend, allocator)` constructor seam. + seam = ContractAlgorithm(DefaultBackend(), ManualAllocator()) @test contract(a1, labels1, a2, labels2; alg = seam)[1] ≈ ref # `nothing` fields fall back to the TensorOperations defaults. From 83af7be47e93c238e2ccb0a2181e03944d9f0e3d Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 07:46:32 -0400 Subject: [PATCH 26/33] Add ndims_domain beside ndims_codomain A type that stores a codomain/domain split can now report both ranks through one pair of accessors. Only `ndims_codomain` needs overloading, since `ndims_domain` is whatever rank is left over. Co-Authored-By: Claude Opus 5 (1M context) --- src/TensorAlgebra.jl | 2 +- src/projectto.jl | 15 +++++++++++++-- test/test_exports.jl | 3 ++- 3 files changed, 16 insertions(+), 4 deletions(-) diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index d2ba8c7c..63968afb 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -5,7 +5,7 @@ export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public AbstractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermalign, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" + "public AbstractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermalign, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, ndims_domain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/src/projectto.jl b/src/projectto.jl index 60a3418a..59c53862 100644 --- a/src/projectto.jl +++ b/src/projectto.jl @@ -72,9 +72,20 @@ end # The flat all-codomain (state) form: a list of `axes` with an empty domain. unchecked_project(raw, axes) = unchecked_project(raw, axes, ()) -# The codomain rank a destination reports when no split is given: its full rank by default (no -# domain), overloaded by a backend that stores a split (a `TensorMap` returns `numout`). +""" + TensorAlgebra.ndims_codomain(a) -> Int + TensorAlgebra.ndims_domain(a) -> Int + +The codomain and domain ranks of `a`'s intrinsic split, for a type that stores one. An array has +no intrinsic split, so it defaults to all codomain and an empty domain, and a type that does store +one overloads `ndims_codomain` (a `TensorMap` returns `numout`). + +Only `ndims_codomain` needs overloading: `ndims_domain` is whatever rank is left over, so the two +agree by construction. +""" ndims_codomain(a) = ndims(a) +@doc (@doc ndims_codomain) +ndims_domain(a) = ndims(a) - ndims_codomain(a) """ is_projected(dest, src, ndims_codomain::Val; kwargs...) -> Bool diff --git a/test/test_exports.jl b/test/test_exports.jl index e945e1c9..51759893 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -33,7 +33,8 @@ using Test: @test, @testset :left_polar, :LinearBroadcasted, :linearbroadcasted, :lq_compact, :lq_full, :matricize, :MatricizeContract, :matricizeop, :matricizeop!, :matricizeopcopy, :matricizeopview, :MatricizeStyle, :MATRIX_FUNCTIONS, - :ndims, :ndims_codomain, :one, :ones_map, :operation, :output_axes, + :ndims, :ndims_codomain, :ndims_domain, :one, :ones_map, :operation, + :output_axes, :PermutedDims, :permuteddims, :permutedims, :permutedims!, :permutedimsadd!, :permutedimsop, :permutedimsopadd!, :project, :project!, :project_aux, :project_hermitian, :projectto!, :qr_compact, From 931c575544184692f04eef870b4846fe0b93cfa1 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 10:41:05 -0400 Subject: [PATCH 27/33] Fill the identity through MatrixAlgebra.one! `TensorAlgebra.one` bottomed out on `MatrixAlgebraKit.one!`, which a `TensorMap` has no method for, so the TensorKit extension worked around it one level up and `one!` on a `TensorMap` never worked at all. The matrix-level fill is now its own entry in `MatrixAlgebra` for a backend to overload, which also lets `one!!` go. Co-Authored-By: Claude Opus 5 (1M context) --- ext/TensorAlgebraTensorKitExt.jl | 4 ++ src/MatrixAlgebra.jl | 11 ++++++ src/factorizations.jl | 65 +++++++++++++------------------- test/test_exports.jl | 1 + test/test_tensorkitext.jl | 18 ++++++++- 5 files changed, 59 insertions(+), 40 deletions(-) diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index a9114ceb..4e0a8f74 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -278,6 +278,10 @@ function TensorAlgebra.matricizeop!( ) end +# TensorKit owns its own `one!` generic and forwards only the `AbstractMatrix` case to +# MatrixAlgebraKit, so the identity fill on a regrouped `TensorMap` goes through TensorKit. +TensorAlgebra.MatrixAlgebra.one!(t::AbstractTensorMap) = TensorKit.one!(t) + function TensorAlgebra.unmatricize( ::TensorKitMatricize, m::AbstractTensorMap, axes_codomain, axes_domain ) diff --git a/src/MatrixAlgebra.jl b/src/MatrixAlgebra.jl index f519acc0..18cc1dd8 100644 --- a/src/MatrixAlgebra.jl +++ b/src/MatrixAlgebra.jl @@ -2,6 +2,7 @@ module MatrixAlgebra export invsqrt_diag_safe, invsqrth_safe, + one!, pow_diag_safe, pow_diag_safe!, powh_safe, @@ -12,6 +13,16 @@ export invsqrt_diag_safe, using LinearAlgebra: LinearAlgebra, Diagonal, isdiag, norm using MatrixAlgebraKit: MatrixAlgebraKit as MAK +""" + MatrixAlgebra.one!(m) -> m + +Fill `m` with the identity in place. The matrix-level identity fill the tensor-level +`TensorAlgebra.one`/`one!` bottom out on, and the customization point a backend overloads when its +matricization is a type `MatrixAlgebraKit.one!` does not handle (TensorKit owns its own `one!` +generic rather than extending MatrixAlgebraKit's, so a `TensorMap` fills through that). +""" +one!(m) = MAK.one!(m) + function _clamp_kwargs_doc(arg::AbstractString) return join( ( diff --git a/src/factorizations.jl b/src/factorizations.jl index 7112d030..40fa3261 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -730,7 +730,7 @@ and is not mutated. A tensor generalization in its own right, not an extension of `Base.one`, so it is neither exported nor imported. Qualify as `TensorAlgebra.one(A, ...)`. -See also `MatrixAlgebraKit.one!`. +See also [`MatrixAlgebra.one!`](@ref), the matrix-level fill this bottoms out on. # Examples @@ -749,66 +749,55 @@ true """ function one end -function one!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) - MatrixAlgebraKit.one!(A_mat) - axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) - return unmatricize(style, A_mat, axes_codomain, axes_domain) -end -function one!!(A, ndims_codomain::Val; kwargs...) - return one!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) -end - # In-place identity fill: writes the identity into `A` and returns it. Fills the memory-sharing # matricization directly when the style declares one at this split, and otherwise fills a # gathered matrix and scatters it back with `unmatricize!`. -function one!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) +function one!(style::MatricizeStyle, A, ndims_codomain::Val) perm_codomain, perm_domain = identitybiperm(ndims_codomain, Val(ndims(A))) if is_output_view(matricizeop, style, identity, A, perm_codomain, perm_domain) - MatrixAlgebraKit.one!( + MatrixAlgebra.one!( matricizeopview(style, identity, A, perm_codomain, perm_domain) ) return A end A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) - MatrixAlgebraKit.one!(A_mat) + MatrixAlgebra.one!(A_mat) return unmatricize!(style, A, A_mat, ndims_codomain) end -function one!(A, ndims_codomain::Val; kwargs...) - return one!(MatricizeStyle(A), A, ndims_codomain; kwargs...) +function one!(A, ndims_codomain::Val) + return one!(MatricizeStyle(A), A, ndims_codomain) end -function one(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - return one!!(style, copy(A), ndims_codomain; kwargs...) +# The fill writes into the matricization, so the bipermutation form is the primitive: +# `matricizeopcopy` permutes and matricizes in one step and hands back storage this call owns, +# which is what makes `A` a shape prototype rather than something to copy up front. +function one(style::MatricizeStyle, A, perm_codomain, perm_domain) + A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) + MatrixAlgebra.one!(A_mat) + axes_codomain, axes_domain = bipartition_axes( + map(i -> axes(A, i), (perm_codomain..., perm_domain...)), + Val(length(perm_codomain)) + ) + return unmatricize(style, A_mat, axes_codomain, axes_domain) end -function one(A, ndims_codomain::Val; kwargs...) - return one!!(copy(A), ndims_codomain; kwargs...) +function one(A, perm_codomain, perm_domain) + return one(MatricizeStyle(A), A, perm_codomain, perm_domain) end - -# `one` stays off the shared factorization wrappers: `one!!` is its own overload point (a -# `TensorMap` backend fills the identity through TensorKit rather than MatrixAlgebraKit). -function one( - style::MatricizeStyle, A, - perm_codomain, perm_domain; - kwargs... - ) - A_perm = bipermutedims(A, perm_codomain, perm_domain) - return one!!(style, A_perm, Val(length(perm_codomain)); kwargs...) +function one(style::MatricizeStyle, A, ndims_codomain::Val) + return one(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) end -function one( - A, perm_codomain, perm_domain; kwargs... - ) - return one(MatricizeStyle(A), A, perm_codomain, perm_domain; kwargs...) +function one(A, ndims_codomain::Val) + return one(MatricizeStyle(A), A, ndims_codomain) end function one( - style::MatricizeStyle, A, labels_A, labels_codomain, labels_domain; kwargs... + style::MatricizeStyle, A, labels_A, labels_codomain, labels_domain ) perm_codomain, perm_domain = biperm(Tuple.((labels_A, labels_codomain, labels_domain))...) - return one(style, A, perm_codomain, perm_domain; kwargs...) + return one(style, A, perm_codomain, perm_domain) end -function one(A, labels_A, labels_codomain, labels_domain; kwargs...) +function one(A, labels_A, labels_codomain, labels_domain) perm_codomain, perm_domain = biperm(Tuple.((labels_A, labels_codomain, labels_domain))...) - return one(A, perm_codomain, perm_domain; kwargs...) + return one(A, perm_codomain, perm_domain) end diff --git a/test/test_exports.jl b/test/test_exports.jl index 51759893..90a3c2f9 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -73,6 +73,7 @@ using Test: @test, @testset :MatrixAlgebra, :invsqrt_diag_safe, :invsqrth_safe, + :var"one!", :pow_diag_safe, :pow_diag_safe!, :powh_safe, diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index 30b36a27..acf8cf4c 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -4,8 +4,8 @@ using StableRNGs: StableRNG using TensorAlgebra: TensorAlgebra, contract, matricize, project, project_aux, projectto!, rand_map, randn_map, similar_map, tryflattenlinear, tryproject, unchecked_project, unmatricize, zeros_map -using TensorKit: @tensor, AbstractTensorMap, DiagonalTensorMap, Irrep, Rep, SU₂, TensorMap, - U₁, dim, dual, fuse, isomorphism, randn, reduceddim, space, storagetype, ←, ⊗ +using TensorKit: TensorKit, @tensor, AbstractTensorMap, DiagonalTensorMap, Irrep, Rep, SU₂, + TensorMap, U₁, dim, dual, fuse, isomorphism, randn, reduceddim, space, storagetype, ←, ⊗ using Test: @test, @test_throws, @testset # A shared bond contracts when it sits in one operand's domain and the other's codomain, i.e. @@ -329,8 +329,22 @@ using Test: @test, @test_throws, @testset # `tr` over a codomain/domain bipartition matches TensorKit's native trace of the endomorphism. t = randn(rng, elt, W ⊗ X, W ⊗ X) + t_before = copy(t) @test TensorAlgebra.tr(t, (:i, :j, :ip, :jp), (:i, :j), (:ip, :jp)) ≈ LinearAlgebra.tr(t) + + # The identity fill routes through `one_matrix!`, which TensorKit answers with its own + # `one!` rather than MatrixAlgebraKit's (that one speaks `AbstractMatrix` only). + Id = TensorAlgebra.one(t, Val(2)) + @test Id ≈ TensorKit.id(TensorKit.domain(t)) + @test Id !== t + @test t ≈ t_before + Id_labels = TensorAlgebra.one(t, (:i, :j, :ip, :jp), (:i, :j), (:ip, :jp)) + @test Id_labels ≈ Id + # In-place, the matching split is TensorKit's own space, so it fills `t` through the view. + c = copy(t) + @test TensorAlgebra.one!(c, Val(2)) === c + @test c ≈ Id end end From fe77d839bd55b8d2765df65257d5fde9d4b9a71f Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 10:59:34 -0400 Subject: [PATCH 28/33] Key the Diagonal contract allocation on the destination shape The arity check looked at the operand bipermutations, which are identical for a `{1,1}` destination and for one that groups both free legs on the same side, so a `{2,0}` destination got a `Diagonal` it cannot represent. Co-Authored-By: Claude Opus 5 (1M context) --- src/diagonal.jl | 20 ++++++++------------ test/test_diagonal.jl | 5 +++++ 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/src/diagonal.jl b/src/diagonal.jl index d4befc6e..c5e3b4a4 100644 --- a/src/diagonal.jl +++ b/src/diagonal.jl @@ -69,31 +69,27 @@ function unmatricize( return unmatricize(style, copyto!(similar(m, axes(m)), m), axes_codomain, axes_domain) end +# Contracting two `Diagonal`s to a `{1,1}` destination is the matmul/endomorphism pattern +# `Diagonal * Diagonal = Diagonal` (all transpose variants `[i,j]*[j,k]`, `[i,j]*[k,j]`, ...), so +# allocate a `Diagonal`. Every other destination shape (rank-4 outer product, scalar full +# contraction, or both free legs grouped on one side) is not representable as a `Diagonal` and does +# not match this signature, falling back to the generic dense allocation the way `Diagonal`/dense +# mixing does. function allocate_output( ::typeof(contract), - perm_dest_codomain, perm_dest_domain, + perm_dest_codomain::Tuple{Int}, perm_dest_domain::Tuple{Int}, a1::Diagonal, perm1_codomain, perm1_domain, a2::Diagonal, perm2_codomain, perm2_domain ) check_input( contract, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) - axes_codomain_dest, axes_domain_dest = output_axes( + axes_codomain_dest, _ = output_axes( contract, perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain ) T = Base.promote_op(matprod, eltype(a1), eltype(a2)) - # Contracting two `Diagonal`s over a single leg, leaving one free leg on each, is the - # matmul/endomorphism pattern `Diagonal * Diagonal = Diagonal` (all transpose variants - # `[i,j]*[j,k]`, `[i,j]*[k,j]`, ...), whose `{1,1}` output stays `Diagonal`. Every other - # pattern (rank-4 outer product, scalar full contraction) is not representable as a - # `Diagonal` and takes the generic dense allocation, matching `Diagonal`/dense mixing. - is_matmul = - length(perm1_codomain) == 1 && length(perm1_domain) == 1 && - length(perm2_domain) == 1 - is_matmul || - return zero!(similar_map(a1, T, axes_codomain_dest, axes_domain_dest)) return Diagonal(zero!(similar(a1.diag, T, (only(axes_codomain_dest),)))) end diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index 45fcefe8..9c7467a8 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -99,6 +99,11 @@ using Test: @test, @test_throws, @testset c0, = TensorAlgebra.contract(d, ("i", "j"), d2, ("i", "j")) @test ndims(c0) == 0 @test c0[] ≈ sum(diag(d) .* diag(d2)) + # A destination bipermutation that groups both free legs on one side is a `{2,0}` map, + # not representable as a `Diagonal`, so it densifies rather than erroring. + cg = TensorAlgebra.contractpermalign((1, 2), (), d, (1,), (2,), d2, (1,), (2,)) + @test !(cg isa Diagonal) + @test cg ≈ d * d2 # No contracted legs: a rank-4 outer product, densified. c4, = TensorAlgebra.contract(d, ("i", "j"), d2, ("k", "l")) @test !(c4 isa Diagonal) From fea40b50cf19a128ad0f50bf7d23596e1d68db14 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 11:03:53 -0400 Subject: [PATCH 29/33] Return one in the same array type as its argument `one` built its result from `matricizeopcopy`, whose dense allocation flattens a `Diagonal` into a `Matrix`. Filling a permuted copy through `one!` instead costs the same single copy and keeps the structure. Nothing asserted the returned type, so the tests now do. Co-Authored-By: Claude Opus 5 (1M context) --- src/factorizations.jl | 16 ++++++---------- test/test_diagonal.jl | 25 +++++++++++++++++++++++++ test/test_factorizations.jl | 10 +++++++++- test/test_tensorkitext.jl | 21 +++++++++++++++++---- 4 files changed, 57 insertions(+), 15 deletions(-) diff --git a/src/factorizations.jl b/src/factorizations.jl index 40fa3261..be1dd13b 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -768,17 +768,13 @@ function one!(A, ndims_codomain::Val) return one!(MatricizeStyle(A), A, ndims_codomain) end -# The fill writes into the matricization, so the bipermutation form is the primitive: -# `matricizeopcopy` permutes and matricizes in one step and hands back storage this call owns, -# which is what makes `A` a shape prototype rather than something to copy up front. +# Fills a permuted copy in place rather than building the matrix itself, so the result keeps the +# structure `bipermutedims` gives it (the identity of a `Diagonal` is a `Diagonal`, which the +# dense `allocate_output` behind `matricizeopcopy` would flatten). The copy is the same one the +# caller would otherwise pay for `A` being a shape prototype. function one(style::MatricizeStyle, A, perm_codomain, perm_domain) - A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) - MatrixAlgebra.one!(A_mat) - axes_codomain, axes_domain = bipartition_axes( - map(i -> axes(A, i), (perm_codomain..., perm_domain...)), - Val(length(perm_codomain)) - ) - return unmatricize(style, A_mat, axes_codomain, axes_domain) + A_perm = bipermutedims(A, perm_codomain, perm_domain) + return one!(style, A_perm, Val(length(perm_codomain))) end function one(A, perm_codomain, perm_domain) return one(MatricizeStyle(A), A, perm_codomain, perm_domain) diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index 9c7467a8..48c50b82 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -79,6 +79,31 @@ using Test: @test, @test_throws, @testset @test e ≈ exp(dp) end + @testset "one and one! preserve Diagonal" begin + Id = Diagonal(ones(elt, 3)) + style = TensorAlgebra.ReshapeMatricize() + for got in ( + TensorAlgebra.one(d, ("i", "j"), ("i",), ("j",)), + TensorAlgebra.one(d, Val(1)), + TensorAlgebra.one(d, (1,), (2,)), + TensorAlgebra.one(d, (2,), (1,)), + TensorAlgebra.one(style, d, Val(1)), + TensorAlgebra.one(style, d, (1,), (2,)), + ) + @test got isa Diagonal + @test got == Id + end + # The allocating forms treat `d` as a shape prototype and leave it alone. + @test d == Diagonal(elt[2, 3, 4]) + + dfill = Diagonal(elt[5, 6, 7]) + @test TensorAlgebra.one!(dfill, Val(1)) === dfill + @test dfill == Id + dstyle = Diagonal(elt[5, 6, 7]) + @test TensorAlgebra.one!(style, dstyle, Val(1)) === dstyle + @test dstyle == Id + end + @testset "contract stays Diagonal on the matmul pattern, densifies otherwise" begin d2 = Diagonal(elt[10, 20, 30]) # One contracted leg: the matmul/endomorphism pattern stays Diagonal. diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index c2daeaf2..b6bfe7db 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -339,12 +339,16 @@ end @test size(Id) == size(A) @test eltype(Id) === T + @test typeof(Id) === typeof(A) @test TensorAlgebra.matricize(Id, splitperms(Id, 2)...) ≈ I - # `Val`, perm, and label entries agree. + # `Val`, perm, and label entries agree, as do the style-explicit spellings of each. + style = TensorAlgebra.MatricizeStyle(A) @test TensorAlgebra.one(A, Val(2)) ≈ Id @test TensorAlgebra.one(A, (1, 2), (3, 4)) ≈ Id + @test TensorAlgebra.one(style, A, Val(2)) ≈ Id + @test TensorAlgebra.one(style, A, (1, 2), (3, 4)) ≈ Id # Non-trivial codomain/domain partition: codomain (a, b) interleaved with # domain (c, d) in the input layout. The result is permuted into the @@ -363,6 +367,10 @@ end @test Cret === C @test TensorAlgebra.matricize(C, splitperms(C, 2)...) ≈ I @test C ≈ TensorAlgebra.one(A, Val(2)) + Cstyle = randn(T, 2, 3, 2, 3) + @test TensorAlgebra.one!(TensorAlgebra.MatricizeStyle(Cstyle), Cstyle, Val(2)) === + Cstyle + @test TensorAlgebra.matricize(Cstyle, splitperms(Cstyle, 2)...) ≈ I # `unmatricize!` scatters a fused matrix back into an existing array. D = randn(T, 2, 3, 2, 3) diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index acf8cf4c..11a2479e 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -333,18 +333,31 @@ using Test: @test, @test_throws, @testset @test TensorAlgebra.tr(t, (:i, :j, :ip, :jp), (:i, :j), (:ip, :jp)) ≈ LinearAlgebra.tr(t) - # The identity fill routes through `one_matrix!`, which TensorKit answers with its own - # `one!` rather than MatrixAlgebraKit's (that one speaks `AbstractMatrix` only). + # The identity fill routes through `MatrixAlgebra.one!`, which TensorKit answers with its + # own `one!` rather than MatrixAlgebraKit's (that one speaks `AbstractMatrix` only). + style = TensorAlgebra.MatricizeStyle(t) Id = TensorAlgebra.one(t, Val(2)) @test Id ≈ TensorKit.id(TensorKit.domain(t)) @test Id !== t + @test space(Id) == space(t) + @test t ≈ t_before + for got in ( + TensorAlgebra.one(t, (:i, :j, :ip, :jp), (:i, :j), (:ip, :jp)), + TensorAlgebra.one(t, (1, 2), (3, 4)), + TensorAlgebra.one(style, t, Val(2)), + TensorAlgebra.one(style, t, (1, 2), (3, 4)), + ) + @test space(got) == space(t) + @test got ≈ Id + end @test t ≈ t_before - Id_labels = TensorAlgebra.one(t, (:i, :j, :ip, :jp), (:i, :j), (:ip, :jp)) - @test Id_labels ≈ Id # In-place, the matching split is TensorKit's own space, so it fills `t` through the view. c = copy(t) @test TensorAlgebra.one!(c, Val(2)) === c @test c ≈ Id + c_style = copy(t) + @test TensorAlgebra.one!(style, c_style, Val(2)) === c_style + @test c_style ≈ Id end end From c36b040b5c94edff67fe8695ae6bd30e89c6a603 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 11:35:00 -0400 Subject: [PATCH 30/33] Test the contract bipermutation ladder directly The labels entry points only ever hand down the canonical destination split, so the unevenly split and permuted destinations the `perm` signatures accept had no coverage. That is where the `Diagonal` allocation bug lived. Co-Authored-By: Claude Opus 5 (1M context) --- test/test_contractperm.jl | 106 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 106 insertions(+) create mode 100644 test/test_contractperm.jl diff --git a/test/test_contractperm.jl b/test/test_contractperm.jl new file mode 100644 index 00000000..61b88622 --- /dev/null +++ b/test/test_contractperm.jl @@ -0,0 +1,106 @@ +using StableRNGs: StableRNG +using TensorAlgebra: TensorAlgebra, MatricizeContract, contractperm, contractperm!, + contractpermadd!, contractpermalign, contractpermopadd! +using Test: @test, @test_throws, @testset + +# The bipermutation rungs of the contract ladder, called directly. The labels entry points only +# ever hand down the canonical `(1:k), (k+1:n)` destination split, so an unevenly split or +# permuted destination is reachable only through these signatures. +@testset "contract bipermutation ladder (eltype=$elt)" for elt in (Float64, ComplexF64) + rng = StableRNG(1234) + a1 = randn(rng, elt, 2, 3, 4) + a2 = randn(rng, elt, 4, 5) + p1 = ((1, 2), (3,)) + p2 = ((1,), (2,)) + # The free legs in their natural order: `a1`'s codomain, then `a2`'s domain. + ref = reshape(reshape(a1, 6, 4) * a2, 2, 3, 5) + # Destination bipermutations, starting with the canonical split the labels forms produce and + # then the ones they cannot: every free leg on one side, and permuted leg orders. + dests = ( + ((1, 2), (3,)), + ((1, 2, 3), ()), + ((), (1, 2, 3)), + ((3, 1), (2,)), + ((2,), (3, 1)), + ((1,), (2, 3)), + ) + destsize(perm) = map(i -> size(ref, i), perm) + + @testset "contractperm takes the canonical destination split" begin + @test contractperm(a1, p1..., a2, p2...) ≈ ref + end + + @testset "contractpermalign honors an arbitrary destination bipermutation" begin + for (dest_codomain, dest_domain) in dests + perm = (dest_codomain..., dest_domain...) + got = contractpermalign(dest_codomain, dest_domain, a1, p1..., a2, p2...) + @test size(got) == destsize(perm) + @test got ≈ permutedims(ref, perm) + end + @test a1 == randn(StableRNG(1234), elt, 2, 3, 4) # operands untouched + end + + @testset "contractperm! writes into the destination" begin + for (dest_codomain, dest_domain) in dests + perm = (dest_codomain..., dest_domain...) + dest = fill(elt(7), destsize(perm)) + @test contractperm!( + dest, dest_codomain, dest_domain, a1, p1..., a2, p2... + ) === dest + @test dest ≈ permutedims(ref, perm) + end + end + + @testset "contractpermadd! scales the contraction and the destination" begin + dest_codomain, dest_domain = (3, 1), (2,) + perm = (dest_codomain..., dest_domain...) + dest = randn(rng, elt, destsize(perm)) + dest_before = copy(dest) + α, β = elt(2), elt(3) + @test contractpermadd!( + dest, dest_codomain, dest_domain, a1, p1..., a2, p2..., α, β + ) === dest + @test dest ≈ α * permutedims(ref, perm) + β * dest_before + end + + @testset "contractpermopadd! applies each operand's op" begin + conj_ref = reshape(reshape(conj(a1), 6, 4) * conj(a2), 2, 3, 5) + dest = fill(elt(7), 2, 3, 5) + @test contractpermopadd!( + dest, (1, 2), (3,), conj, a1, p1..., conj, a2, p2..., true, false + ) === dest + @test dest ≈ conj_ref + # Only the first operand conjugated, on a permuted destination. + half_ref = permutedims(reshape(reshape(conj(a1), 6, 4) * a2, 2, 3, 5), (3, 1, 2)) + dest_half = fill(elt(7), destsize((3, 1, 2))) + contractpermopadd!( + dest_half, (3, 1), (2,), conj, a1, p1..., identity, a2, p2..., true, false + ) + @test dest_half ≈ half_ref + end + + @testset "malformed destinations and algorithms are rejected" begin + @test_throws ArgumentError contractpermalign( + (1, 1), (3,), a1, p1..., a2, p2... + ) + @test_throws DimensionMismatch contractperm!( + zeros(elt, 2, 3, 6), (1, 2), (3,), a1, p1..., a2, p2... + ) + @test_throws ArgumentError contractpermalign( + (1, 2), (3,), a1, p1..., randn(rng, elt, 3, 5), p2... + ) + @test_throws ArgumentError contractpermalign( + (1, 2), (3,), a1, p1..., a2, p2...; alg = :not_an_algorithm + ) + end + + @testset "a named algorithm reaches the same result" begin + dest_codomain, dest_domain = (3, 1), (2,) + expected = permutedims(ref, (dest_codomain..., dest_domain...)) + for alg in (TensorAlgebra.DefaultContractAlgorithm(), MatricizeContract()) + @test contractpermalign( + dest_codomain, dest_domain, a1, p1..., a2, p2...; alg + ) ≈ expected + end + end +end From d392c8cbc1b218164ef0c1f53e9f4595a903c3c3 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 11:59:22 -0400 Subject: [PATCH 31/33] Allocate a conjugated TensorMap matricization over the dualized space `matricize` with `conj` threw a `SpaceMismatch` on every `TensorMap`. The destination was allocated over the permuted space rather than the conjugated one, and TensorKit conjugates by adjointing its source, so the two never lined up. Co-Authored-By: Claude Opus 5 (1M context) --- ext/TensorAlgebraTensorKitExt.jl | 12 +-- test/test_diagonal.jl | 36 +++++++++ test/test_tensorkitext.jl | 129 ++++++++++++++++++++++++++++++- 3 files changed, 169 insertions(+), 8 deletions(-) diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 4e0a8f74..a17957e1 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -259,15 +259,17 @@ function TensorAlgebra.matricizeopview( ) return t end -# A `TensorMap`'s matricization is a regrouping of its indices, so the destination is a `TensorMap` -# over the regrouped space and the write is the ordinary permuted-add. `bipermutedimsopadd!` above -# routes that through `tensoradd!`, which realizes the permutation, the `op === conj` conjugation -# and the scaling in one call, so no separate handling of `op` is needed here. +# A `TensorMap`'s matricization is a regrouping of its indices, so the destination is the same +# `TensorMap` the plain permuted-add would allocate. That one already dualizes each space under +# `op === conj`, which is what `bipermutedimsopadd!` needs: `tensoradd!` realizes a conjugation +# by adjointing its source, so a destination over the undualized space does not match it. function TensorAlgebra.allocate_output( ::typeof(TensorAlgebra.matricizeop), ::TensorKitMatricize, op, t::AbstractTensorMap, perm_codomain, perm_domain ) - return similar(t, permute(space(t), (perm_codomain, perm_domain))) + return TensorAlgebra.allocate_output( + TensorAlgebra.permutedimsop, op, t, perm_codomain, perm_domain + ) end function TensorAlgebra.matricizeop!( dest::AbstractTensorMap, ::TensorKitMatricize, op, diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index 48c50b82..fd751107 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -79,6 +79,42 @@ using Test: @test, @test_throws, @testset @test e ≈ exp(dp) end + @testset "the matricize hooks cohere on the aliasing split" begin + style = TensorAlgebra.MatricizeStyle(d) + dz = Diagonal(elt <: Complex ? elt[1 + 2im, 3 - im, 2im] : elt[1, 3, 2]) + for op in (identity, conj), (pc, pd) in (((1,), (2,)), ((2,), (1,))) + ref = TensorAlgebra.matricizeop(style, op, dz, pc, pd) + if TensorAlgebra.is_output_view( + TensorAlgebra.matricizeop, + style, + op, + dz, + pc, + pd + ) + m = TensorAlgebra.matricizeopview(style, op, dz, pc, pd) + @test Base.mightalias(m, dz) + @test m == ref + end + m_copy = TensorAlgebra.matricizeopcopy(style, op, dz, pc, pd) + @test !Base.mightalias(m_copy, dz) + @test m_copy == ref + end + # Only the identity op on the untransposed split aliases. `matricizeopview` hands back + # `dz` itself, so a declared share under `conj` would silently skip the conjugation. + @test TensorAlgebra.is_output_view( + TensorAlgebra.matricizeop, style, identity, dz, (1,), (2,) + ) + @test !TensorAlgebra.is_output_view( + TensorAlgebra.matricizeop, style, conj, dz, (1,), (2,) + ) + + dest = Diagonal(zeros(elt, 3)) + src = Diagonal(elt[7, 8, 9]) + @test TensorAlgebra.unmatricize!(style, dest, src, Val(1)) === dest + @test dest == src + end + @testset "one and one! preserve Diagonal" begin Id = Diagonal(ones(elt, 3)) style = TensorAlgebra.ReshapeMatricize() diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index 11a2479e..e4446fdc 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -1,9 +1,11 @@ using Base.Broadcast: broadcasted using LinearAlgebra: LinearAlgebra, norm using StableRNGs: StableRNG -using TensorAlgebra: TensorAlgebra, contract, matricize, project, project_aux, projectto!, - rand_map, randn_map, similar_map, tryflattenlinear, tryproject, unchecked_project, - unmatricize, zeros_map +using TensorAlgebra: TensorAlgebra, contract, contractalign, eig_full, eig_vals, eigh_full, + eigh_vals, left_null, left_orth, left_polar, lq_compact, lq_full, matricize, project, + project_aux, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, + right_orth, right_polar, similar_map, svd_compact, svd_full, svd_vals, tryflattenlinear, + tryproject, unchecked_project, unmatricize, zeros_map using TensorKit: TensorKit, @tensor, AbstractTensorMap, DiagonalTensorMap, Irrep, Rep, SU₂, TensorMap, U₁, dim, dual, fuse, isomorphism, randn, reduceddim, space, storagetype, ←, ⊗ using Test: @test, @test_throws, @testset @@ -359,6 +361,127 @@ using Test: @test, @test_throws, @testset @test TensorAlgebra.one!(style, c_style, Val(2)) === c_style @test c_style ≈ Id end + + @testset "the matricize hooks cohere on the matching split" begin + W = Rep[U₁](0 => 2, 1 => 1) + X = Rep[U₁](0 => 1, 1 => 2) + Z = Rep[U₁](0 => 1, 1 => 2) + t = randn(rng, elt, W ⊗ X, Z) + style = TensorAlgebra.MatricizeStyle(t) + splits = (((1, 2), (3,)), ((1, 3), (2,)), ((1, 2, 3), ()), ((), (1, 2, 3))) + for op in (identity, conj), (pc, pd) in splits + # Regrouping a `TensorMap` is a permuted-add, so the matricization is the tensor + # `permutedimsop` builds, spaces included. Under `conj` that dualizes every space. + ref = TensorAlgebra.permutedimsop(op, t, pc, pd) + if TensorAlgebra.is_output_view(TensorAlgebra.matricizeop, style, op, t, pc, pd) + @test op === identity # only the identity op can hand back `t` itself + @test TensorAlgebra.matricizeopview(style, op, t, pc, pd) === t + end + m_copy = TensorAlgebra.matricizeopcopy(style, op, t, pc, pd) + @test m_copy !== t + @test TensorAlgebra.data(m_copy) !== TensorAlgebra.data(t) + @test space(m_copy) == space(ref) + @test m_copy ≈ ref + end + # Only the split TensorKit already stores is a view. A regrouped one has to be built. + @test TensorAlgebra.is_output_view( + TensorAlgebra.matricizeop, style, identity, t, (1, 2), (3,) + ) + @test !TensorAlgebra.is_output_view( + TensorAlgebra.matricizeop, style, identity, t, (1, 3), (2,) + ) + + dest = similar(t) + @test TensorAlgebra.unmatricize!( + style, dest, matricize(style, t, (1, 2), (3,)), Val(2) + ) === dest + @test dest ≈ t + end + + # The factorizations are generic over the matricize hooks, so a `TensorMap` reaches them + # through the extension rather than through any dedicated method. Each factor is checked for + # the spaces it carries as well as for reconstructing the input. + @testset "factorizations" begin + W = Rep[U₁](0 => 2, 1 => 1) + X = Rep[U₁](0 => 1, 1 => 2) + # Every charge sector of the fused codomain is at least as large as the domain's, which + # is what the tall-block factorizations need. + Z = Rep[U₁](0 => 1, 1 => 2) + t = randn(rng, elt, W ⊗ X, Z) + t_before = copy(t) + labels_t = (:i, :j, :k) + labels_l = (:i, :j) + labels_r = (:k,) + # Both halves of every two-factor form carry the bond on the inside, so one + # reconstruction covers them all. + reconstruct(F1, F2) = contractalign( + labels_t, F1, (labels_l..., :q), F2, (:q, labels_r...) + ) + + @testset "$(nameof(f))" for f in ( + qr_compact, qr_full, lq_compact, lq_full, left_orth, right_orth, left_polar, + ) + F1, F2 = f(t, labels_t, labels_l, labels_r) + @test F1 isa AbstractTensorMap + @test F2 isa AbstractTensorMap + @test space(F1, 1) == space(t, 1) + @test space(F1, 2) == space(t, 2) + @test space(F2, 2) == space(t, 3) + @test reconstruct(F1, F2) ≈ t + end + + @testset "$(nameof(f))" for f in (svd_compact, svd_full) + U, S, Vᴴ = f(t, labels_t, labels_l, labels_r) + @test U isa AbstractTensorMap + @test Vᴴ isa AbstractTensorMap + US, labels_US = contract(U, (labels_l..., :u), S, (:u, :q)) + @test contractalign(labels_t, US, labels_US, Vᴴ, (:q, labels_r...)) ≈ t + end + + @testset "svd_vals matches the compact spectrum" begin + _, S, _ = svd_compact(t, labels_t, labels_l, labels_r) + @test svd_vals(t, labels_t, labels_l, labels_r) ≈ TensorAlgebra.data(S) + end + + @testset "the left null space annihilates the map" begin + N = left_null(t, labels_t, labels_l, labels_r) + @test N isa AbstractTensorMap + @test norm(N' * t) < sqrt(eps(real(elt))) * norm(t) + end + + @test t ≈ t_before # no factorization altered the input + + # `right_polar` needs a wide block in every sector, and `t`'s right null space is empty + # because it has full column rank, so both get the transposed map. + @testset "the wide map: right_polar and the right null space" begin + tw = randn(rng, elt, Z, W ⊗ X) + labels_w = (:k, :i, :j) + P, Qw = right_polar(tw, labels_w, (:k,), (:i, :j)) + @test contractalign(labels_w, P, (:k, :q), Qw, (:q, :i, :j)) ≈ tw + M = right_null(tw, labels_w, (:k,), (:i, :j)) + @test M isa AbstractTensorMap + @test norm(tw * M') < sqrt(eps(real(elt))) * norm(tw) + end + + @testset "the eigen family on an endomorphism" begin + te = randn(rng, elt, W ⊗ X, W ⊗ X) + labels_e = (:i, :j, :ip, :jp) + labels_v, labels_vp = (:i, :j), (:ip, :jp) + D, V = eig_full(te, labels_e, labels_v, labels_vp) + teV = contractalign((:i, :j, :d), te, labels_e, V, (labels_vp..., :d)) + VD = contractalign((:i, :j, :d), V, (labels_v..., :dp), D, (:dp, :d)) + @test teV ≈ VD + @test eig_vals(te, labels_e, labels_v, labels_vp) ≈ TensorAlgebra.data(D) + + # The Hermitian entries need a Hermitian input. + th = te + te' + Dh, Vh = eigh_full(th, labels_e, labels_v, labels_vp) + thV = contractalign((:i, :j, :d), th, labels_e, Vh, (labels_vp..., :d)) + VhD = contractalign((:i, :j, :d), Vh, (labels_v..., :dp), Dh, (:dp, :d)) + @test thV ≈ VhD + @test eigh_vals(th, labels_e, labels_v, labels_vp) ≈ TensorAlgebra.data(Dh) + end + end end @testset "dual/isdual on TensorKit spaces and sectors" begin From 5e6704ebc431bf5495e29d4a43bc746650af0116 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 12:04:17 -0400 Subject: [PATCH 32/33] Check that every declared public name resolves `public` accepts a name that resolves to nothing, and the expected-name list is maintained alongside the declaration, so a typo gets edited into both and the set comparison still passes. Co-Authored-By: Claude Opus 5 (1M context) --- test/test_exports.jl | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/test/test_exports.jl b/test/test_exports.jl index 90a3c2f9..51e8419f 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -50,6 +50,10 @@ using Test: @test, @testset ) end @test issetequal(names(TensorAlgebra), exports) + # `public` accepts a name that resolves to nothing, and the list above is maintained + # alongside the declaration, so a typo would be edited into both and slip past the + # comparison. Check the declared names actually resolve. + @test all(n -> isdefined(TensorAlgebra, n), names(TensorAlgebra)) # The matrix-level factorizations are `public`, not exported: the names MatrixAlgebraKit # also exports would otherwise collide in a session loading both packages, and the @@ -82,4 +86,8 @@ using Test: @test, @testset :sqrth_safe, ] @test issetequal(names(TensorAlgebra.MatrixAlgebra), exports) + @test all( + n -> isdefined(TensorAlgebra.MatrixAlgebra, n), + names(TensorAlgebra.MatrixAlgebra) + ) end From 166063a4be4a199f80aa460e5edcffce05c35c37 Mon Sep 17 00:00:00 2001 From: Matthew Fishman Date: Thu, 24 Sep 2026 12:07:07 -0400 Subject: [PATCH 33/33] Drop the public-name resolvability check --- test/test_exports.jl | 8 -------- 1 file changed, 8 deletions(-) diff --git a/test/test_exports.jl b/test/test_exports.jl index 51e8419f..90a3c2f9 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -50,10 +50,6 @@ using Test: @test, @testset ) end @test issetequal(names(TensorAlgebra), exports) - # `public` accepts a name that resolves to nothing, and the list above is maintained - # alongside the declaration, so a typo would be edited into both and slip past the - # comparison. Check the declared names actually resolve. - @test all(n -> isdefined(TensorAlgebra, n), names(TensorAlgebra)) # The matrix-level factorizations are `public`, not exported: the names MatrixAlgebraKit # also exports would otherwise collide in a session loading both packages, and the @@ -86,8 +82,4 @@ using Test: @test, @testset :sqrth_safe, ] @test issetequal(names(TensorAlgebra.MatrixAlgebra), exports) - @test all( - n -> isdefined(TensorAlgebra.MatrixAlgebra, n), - names(TensorAlgebra.MatrixAlgebra) - ) end