diff --git a/Project.toml b/Project.toml index 146f410b..3fe422fc 100644 --- a/Project.toml +++ b/Project.toml @@ -1,13 +1,12 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.20.1" +version = "0.21.0" authors = ["ITensor developers and contributors"] [workspace] 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/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/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl index 713a3e7d..99961560 100644 --- a/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl +++ b/ext/TensorAlgebraMooncakeExt/TensorAlgebraMooncakeExt.jl @@ -3,7 +3,7 @@ module TensorAlgebraMooncakeExt using Mooncake: Mooncake, @zero_derivative, DefaultCtx using TensorAlgebra: BiTuple, ContractAlgorithm, 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 @@ -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} +@zero_derivative DefaultCtx Tuple{typeof(select_algorithm), Any, Any, Any} end diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index aa770db2..a17957e1 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -243,44 +243,47 @@ 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 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 TensorAlgebra.allocate_output( + TensorAlgebra.permutedimsop, op, 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. +# 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 ) @@ -296,7 +299,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 @@ -310,8 +313,9 @@ function TensorAlgebra.add!( return VectorInterface.add!(y, x, α, β) end -function TensorAlgebra.default_contract_algorithm( - ::Type{<:AbstractTensorMap}, ::Type{<:AbstractTensorMap} +function TensorAlgebra.default_algorithm( + ::typeof(TensorAlgebra.contract!), + ::Type{<:Tuple{AbstractTensorMap, AbstractTensorMap, AbstractTensorMap}} ) return TensorAlgebra.ContractAlgorithm(TO.DefaultBackend()) end diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 2b62f80e..0bcb5134 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -1,63 +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) + return TensorOperationsContract(; backend) +end function TA.ContractAlgorithm(backend::TO.AbstractBackend, allocator) - return TensorOperationsAlgorithm(; backend, allocator) + return TensorOperationsContract(; backend, allocator) end # Using TensorOperations backends as TensorAlgebra implementations # ---------------------------------------------------------------- -# not in-place -function TA.contract( - 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.contract( - 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.contractopadd!( - algorithm::TensorOperationsAlgorithm, +function TA.contractpermopadd!( + algorithm::TensorOperationsContract, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, @@ -89,7 +56,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/MatrixAlgebra.jl b/src/MatrixAlgebra.jl index 82aeece6..18cc1dd8 100644 --- a/src/MatrixAlgebra.jl +++ b/src/MatrixAlgebra.jl @@ -1,19 +1,28 @@ module MatrixAlgebra -export gram_eigh_full, - gram_eigh_full_with_pinv, - invsqrt_diag_safe, +export invsqrt_diag_safe, invsqrth_safe, + one!, pow_diag_safe, pow_diag_safe!, powh_safe, sqrt_diag_safe, - sqrth_safe, - sqrth_invsqrth_safe + sqrth_invsqrth_safe, + sqrth_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( ( @@ -193,92 +202,6 @@ function sqrth_invsqrth_safe(M; alg = nothing, kwargs...) 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 2853d1f3..63968afb 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -1,15 +1,11 @@ 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, - 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 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 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 @@ -25,6 +21,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..7a505ad1 --- /dev/null +++ b/src/algorithm.jl @@ -0,0 +1,45 @@ +""" + TensorAlgebra.AbstractAlgorithm + +Supertype for the algorithm objects operations dispatch on. An operation's own supertype +subtypes this (for example `ContractAlgorithm`), which is what makes an instance pass +through [`TensorAlgebra.select_algorithm`](@ref) unchanged. +""" +abstract type AbstractAlgorithm end + +""" + 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 +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::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::Tuple) + return throw(ArgumentError("`$alg` is not an algorithm for `$f`")) +end + +""" + 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, 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!, ::Type{<:Tuple{A_dest, A1, A2}})`. +""" +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/bituple.jl b/src/bituple.jl index 17acf4cb..604cee00 100644 --- a/src/bituple.jl +++ b/src/bituple.jl @@ -70,3 +70,26 @@ 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...)) + +# 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 + +# 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/contract/allocate_output.jl b/src/contract/allocate_output.jl index c09f6716..5fb32103 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 @@ -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..1cf720ed 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -2,15 +2,57 @@ # 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. + +```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 +2-element Vector{Symbol}: + :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( + +""" + 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, +i.e. `issetequal(labels_dest, symdiff(labels1, labels2))` must be `true`. + +```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 +60,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 +71,21 @@ function contract( kwargs... ) end -function _contract( +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 contractpermalign( + biperm_dest..., a1, biperm1..., a2, biperm2...; kwargs... + ) end -# contract (bipartitioned permutations) -function contract( +# contractperm (bipartitioned permutations) +# `perm` marks the whole biperm ladder: every rung has a labels-form sibling under the plain name, +# 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; kwargs... @@ -48,14 +94,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 contractpermalign( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; kwargs... ) end -function contract( +function contractpermalign( perm_dest_codomain, perm_dest_domain, a1, perm1_codomain, perm1_domain, a2, perm2_codomain, perm2_domain; @@ -67,7 +113,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 +132,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 +158,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 +199,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 ) check_input( contract!, @@ -171,8 +217,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!, alg, (a_dest, a1, a2)) + return contractpermopadd!( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, op1, a1, perm1_codomain, perm1_domain, @@ -180,9 +226,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 +237,7 @@ function contractopadd!( ) return throw( MethodError( - contractopadd!, + contractpermopadd!, ( algorithm, a_dest, perm_dest_codomain, perm_dest_domain, @@ -202,3 +248,18 @@ function contractopadd!( ) ) end + +# 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!), ::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, args::Tuple) + return default_algorithm(contract!, args) +end diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 9e04f5f0..7b4e66e3 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -1,7 +1,7 @@ using LinearAlgebra: mul! -function contractopadd!( - algorithm::Matricize, +function contractpermopadd!( + 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, @@ -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 @@ -34,13 +38,15 @@ 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) + a_dest_mat = matricizeopcopy( + output_style, identity, 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/contract/contractalgorithm.jl b/src/contract/contractalgorithm.jl index 02fb0183..e7ffc8a7 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -1,40 +1,27 @@ -abstract type ContractAlgorithm end +abstract type ContractAlgorithm <: AbstractAlgorithm end ContractAlgorithm(algorithm::ContractAlgorithm) = algorithm struct DefaultContractAlgorithm <: ContractAlgorithm end -struct Matricize{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm +struct MatricizeContract{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm 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} <: + ContractAlgorithm backend::Backend = nothing allocator::Allocator = nothing end - -function select_contract_algorithm(algorithm, a1, a2) - return error("Not implemented.") -end -function select_contract_algorithm(algorithm::ContractAlgorithm, a1, a2) - return algorithm -end -function select_contract_algorithm(algorithm::DefaultContractAlgorithm, a1, a2) - return default_contract_algorithm(a1, a2) -end -function default_contract_algorithm(a1, a2) - return default_contract_algorithm(typeof(a1), typeof(a2)) -end -function default_contract_algorithm(A1::Type{<:AbstractArray}, A2::Type{<:AbstractArray}) - return Matricize(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) -end diff --git a/src/diagonal.jl b/src/diagonal.jl index 7a465893..c5e3b4a4 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. @@ -65,13 +69,27 @@ 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} +# 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::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, _ = 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)) + return Diagonal(zero!(similar(a1.diag, T, (only(axes_codomain_dest),)))) end diff --git a/src/factorizations.jl b/src/factorizations.jl index 1ab39631..be1dd13b 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 ( @@ -29,16 +29,21 @@ 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) || 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,18 +61,17 @@ 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, ) @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) + 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...)), @@ -85,16 +89,13 @@ 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, + :left_null, :right_null, :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...) @@ -102,7 +103,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 +144,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 @@ -168,13 +169,15 @@ true ``` """ function tr(style::MatricizeStyle, A, ndims_codomain::Val) - return LinearAlgebra.tr(matricize(style, 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) end -function tr(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) - return LinearAlgebra.tr(matricizeperm(A, perm_codomain, perm_domain)) +function tr(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 = @@ -184,7 +187,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 +205,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 +223,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 +241,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 +259,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 +276,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 +293,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 +310,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 +386,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 +399,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 +412,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 @@ -425,15 +428,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) @@ -446,7 +449,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 +462,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 +475,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 +489,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 +501,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 +513,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 +525,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 +538,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 @@ -554,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, 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))),)) @@ -572,7 +575,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 @@ -591,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, 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) @@ -607,127 +610,9 @@ 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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; 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, 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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; 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, 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::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 +635,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 +669,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 +689,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 +718,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 @@ -845,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 @@ -858,69 +743,57 @@ 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) - 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...) - if ismatricizeview(style, A, ndims_codomain) - MatrixAlgebraKit.one!(matricizeview(style, A, ndims_codomain)) +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) + MatrixAlgebra.one!( + matricizeopview(style, identity, A, perm_codomain, perm_domain) + ) return A end - A_mat = matricizecopy(style, A, ndims_codomain) - MatrixAlgebraKit.one!(A_mat) + A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) + 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...) +# 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_perm = bipermutedims(A, perm_codomain, perm_domain) + return one!(style, A_perm, Val(length(perm_codomain))) 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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; - 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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; 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/src/matricize.jl b/src/matricize.jl index 972cb072..757dad8f 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 @@ -73,160 +70,133 @@ 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. +# +# `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`. -# `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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} - ) - a_perm = bipermutedims(a, perm_codomain, perm_domain) - return matricize(style, a_perm, Val(length(perm_codomain))) +# 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(ndims_codomain, Val(ndims(a)))...) end - -function matricizeperm( - a, - perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} - ) - 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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} - ) - return matricizeopperm(style, identity, a, perm_codomain, perm_domain) +function matricize(style::MatricizeStyle, a, ndims_codomain::Val) + return matricize(style, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) 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]) +function matricizeop(op, a, ndims_codomain::Val) + return matricizeop(op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end -function to_permblocks( - a, permblocks::Tuple{Tuple{Vararg{Int}}, Tuple{Ellipsis}} - ) - permblocks2 = tuplesetcomplement(ntuple(identity, ndims(a)), permblocks[1]) - return (permblocks[1], permblocks2) +function matricizeop(style::MatricizeStyle, op, a, ndims_codomain::Val) + return matricizeop(style, op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end -function matricizeperm(a, perm_codomain, perm_domain) - return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) +# 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 matricizeperm( - style::MatricizeStyle, a, perm_codomain, perm_domain - ) - return matricizeperm(style, a, to_permblocks(a, (perm_codomain, perm_domain))...) +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 -# ================================== 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) - return matricizeopperm(MatricizeStyle(a), op, 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)) + ) end -function matricizeopperm( - style::MatricizeStyle, op, 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, op, a, to_permblocks(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 -# `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::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} +# 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) + ) ) - ndims(a) == length(perm_codomain) + length(perm_domain) || - throw(ArgumentError("Invalid bipermutation")) - 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) - return matricize(style, a_perm_op, Val(length(perm_codomain))) 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, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} +# ================================== 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 ) - isidentityperm((invperm_codomain..., invperm_domain...)) || return false - return ismatricizeview(style, a, Val(length(invperm_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` # (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) @@ -245,65 +215,37 @@ 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!( +# 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 `matricize`/`unmatricize!` round trip passes +# the same forward bipermutation to both. +function unmatricize!( a_dest, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) - return unmatricizeperm!(MatricizeStyle(m), a_dest, m, invperm_codomain, invperm_domain) + return unmatricize!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) end -function unmatricizeperm!( +function unmatricize!( style::MatricizeStyle, a_dest, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) - 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 -# 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!( - style, - a_dest, - m, - ntuple(identity, Val(K)), - ntuple(i -> K + i, Val(N - K)) + return unmatricize!( + style, a_dest, m, identitybiperm(ndims_codomain, Val(ndims(a_dest)))... ) end function unmatricize!(a_dest, m, ndims_codomain::Val) @@ -313,16 +255,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 ebbda116..a128f22f 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -32,16 +32,13 @@ 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 - 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...) @@ -50,10 +47,10 @@ 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) + 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)) @@ -63,7 +60,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/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/Project.toml b/test/Project.toml index 08c7d265..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" @@ -37,7 +35,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..1903aab2 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -1,9 +1,8 @@ import TensorAlgebra -using EllipsisNotation: var".." using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, - contract!, contractadd!, length_codomain, length_domain, matricizeperm, unmatricize, - unmatricizeperm, unmatricizeperm! + contract!, contractadd!, contractalign, length_codomain, length_domain, matricize, + unmatricize, unmatricize! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -54,70 +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, (..,), (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)) - 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 = matricize(a, (), (1, 2, 3, 4)) @test eltype(a_fused) === elt @test a_fused ≈ reshape(a, (1, 120)) - a_fused = matricizeperm(a, (..,), ()) + 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 @@ -130,30 +127,23 @@ 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)) + unmatricize!(a, m, (1, 2), (3, 4)) @test a ≈ a0 - m1 = matricizeperm(a0, perm_codomain, perm_domain) - a = unmatricizeperm(m1, axes0, perm_codomain, perm_domain) + m1 = matricize(a0, perm_codomain, perm_domain) + a = similar(a0) + 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) @@ -174,8 +164,30 @@ 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,)) + @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) + # 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 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_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.MatricizeContract() + ) ≈ a1 * a2 end @testset "contract eltype widens like a matrix product" begin @@ -194,8 +206,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) @@ -234,14 +246,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 @@ -279,7 +291,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. @@ -302,7 +314,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) @@ -449,17 +461,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_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 diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index c492918a..fd751107 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 @@ -79,6 +79,67 @@ 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() + 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. @@ -99,6 +160,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) diff --git a/test/test_exports.jl b/test/test_exports.jl index 5bab6ad0..90a3c2f9 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 @@ -6,59 +7,73 @@ using Test: @test, @testset :TensorAlgebra, :contract, :contract!, + :contractalign, :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, - :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, [ - :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, + :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 @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, - :gram_eigh_full, - :gram_eigh_full_with_pinv, :invsqrt_diag_safe, :invsqrth_safe, + :var"one!", :pow_diag_safe, :pow_diag_safe!, :powh_safe, diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index f8934459..b6bfe7db 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -1,12 +1,14 @@ 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 +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 +# 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 @@ -20,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 @@ -40,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 @@ -56,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 @@ -73,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 @@ -93,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′) @@ -116,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′) @@ -137,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 @@ -172,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) @@ -183,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 @@ -210,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). @@ -227,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 @@ -250,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 @@ -264,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 @@ -278,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 @@ -295,65 +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...)) -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 + @test A ≈ contractalign(labels_A, P, (labels_P..., :w), W, (:w, labels_W...)) end # one (identity tensor) @@ -372,12 +339,16 @@ end @test size(Id) == size(A) @test eltype(Id) === T + @test typeof(Id) === typeof(A) - @test TensorAlgebra.matricize(Id, Val(2)) ≈ I + @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 @@ -386,7 +357,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 +365,16 @@ 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)) + 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) - 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 +405,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 +433,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 +458,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.matricize(::AliasingMatricize, a::AbstractArray, ndims_codomain::Val) - return TA.matricize(TA.ReshapeMatricize(), a, ndims_codomain) + 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.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 +515,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..1b2dd58c 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,92 +14,101 @@ 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. 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)...) @@ -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 c43627f1..5dd25b17 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 @@ -9,19 +10,35 @@ 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.unmatricizeperm!( + function TA.unmatricize!( ::MyArrayMatricize, a_dest::MyArray, m, - invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} + perm_codomain, perm_domain ) - TA.unmatricizeperm!( - TA.ReshapeMatricize(), a_dest.parent, m, invperm_codomain, invperm_domain + TA.unmatricize!( + TA.ReshapeMatricize(), a_dest.parent, m, perm_codomain, perm_domain ) return a_dest end @@ -38,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!, Tuple{typeof(a1), typeof(a1), typeof(a1)}) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a1), typeof(a1), typeof(a2)}) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a2), typeof(a2), typeof(a1)}) ≡ + MatricizeContract(ReshapeMatricize()) + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a2), typeof(a2), typeof(a2)}) ≡ + MatricizeContract(MyArrayMatricize()) end @testset "style threads through the unfold" begin diff --git a/test/test_matrixalgebra.jl b/test/test_matrixalgebra.jl index e279ea4d..9eaea1f7 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) @@ -195,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 diff --git a/test/test_mooncakeext.jl b/test/test_mooncakeext.jl index 840b5f3c..443df37a 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!, default_contract_algorithm, select_contract_algorithm +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 @testset "MooncakeExt" begin @@ -16,7 +16,7 @@ using Test: @test, @testset @test Mooncake.tangent_type(BiTuple) ≡ Mooncake.NoTangent @test Mooncake.tangent_type(ContractAlgorithm) ≡ 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,17 +55,17 @@ 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 @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 ) diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index 557b4983..e4446fdc 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -1,11 +1,13 @@ 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 TensorKit: @tensor, AbstractTensorMap, DiagonalTensorMap, Irrep, Rep, SU₂, TensorMap, - U₁, dim, dual, fuse, isomorphism, randn, reduceddim, space, storagetype, ←, ⊗ +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 # A shared bond contracts when it sits in one operand's domain and the other's codomain, i.e. @@ -60,7 +62,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 @@ -329,8 +331,156 @@ 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 `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 + # 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 + + @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 diff --git a/test/test_tensoroperations.jl b/test/test_tensoroperations.jl index 68f9c9ee..3f025443 100644 --- a/test/test_tensoroperations.jl +++ b/test/test_tensoroperations.jl @@ -1,5 +1,5 @@ using TensorAlgebra: - ContractAlgorithm, Matricize, TensorOperationsAlgorithm, contract, contract! + ContractAlgorithm, 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 ContractAlgorithm @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 @@ -155,5 +155,5 @@ end @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