diff --git a/Project.toml b/Project.toml index d2b09ee2..db4a3287 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.18.0" +version = "0.19.0" authors = ["ITensor developers and contributors"] [workspace] diff --git a/docs/Project.toml b/docs/Project.toml index 15e0ac72..fbf8bf51 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.18" +TensorAlgebra = "0.19" diff --git a/examples/Project.toml b/examples/Project.toml index dbaa6b09..d560cc23 100644 --- a/examples/Project.toml +++ b/examples/Project.toml @@ -5,4 +5,4 @@ TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" path = ".." [compat] -TensorAlgebra = "0.18" +TensorAlgebra = "0.19" diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index ff9270ec..6b73a832 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -240,11 +240,11 @@ end # A `TensorMap` is already a linear map codomain ← domain, so "matricizing" is just regrouping # its indices into the requested codomain/domain bipartition (`permute`). No fusion or copy of # the array vocabulary is needed: MatrixAlgebraKit factorizes the regrouped `TensorMap` directly. -struct TensorKitFusion <: TensorAlgebra.FusionStyle end -TensorAlgebra.FusionStyle(::Type{<:AbstractTensorMap}) = TensorKitFusion() +struct TensorKitMatricize <: TensorAlgebra.MatricizeStyle end +TensorAlgebra.MatricizeStyle(::Type{<:AbstractTensorMap}) = TensorKitMatricize() function TensorAlgebra.matricize( - ::TensorKitFusion, t::AbstractTensorMap, ndims_codomain::Val{K} + ::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)))) @@ -253,7 +253,7 @@ end # The identity fill on the regrouped map is TensorKit's own `one!` (MatrixAlgebraKit's # `one!` speaks `AbstractMatrix` only). function TensorAlgebra.one!!( - style::TensorKitFusion, A::AbstractTensorMap, ndims_codomain::Val; kwargs... + style::TensorKitMatricize, A::AbstractTensorMap, ndims_codomain::Val; kwargs... ) return TensorKit.one!(TensorAlgebra.matricize(style, A, ndims_codomain)) end @@ -264,7 +264,7 @@ end # codomain-facing (un-dualized), which is exactly TensorKit's domain convention, so they build the # domain `ProductSpace` directly. function TensorAlgebra.unmatricize( - ::TensorKitFusion, m::AbstractTensorMap, codomain_axes, domain_axes + ::TensorKitMatricize, m::AbstractTensorMap, codomain_axes, domain_axes ) S = spacetype(m) dest = ProductSpace{S}(codomain_axes...) ← ProductSpace{S}(domain_axes...) diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 23d173d9..3ade05ab 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -17,12 +17,12 @@ function contractopadd!( a2, biperm2_codomain, biperm2_domain ) a1_mat = matricizeopperm( - algorithm.left_fusion_style, op1, a1, biperm1_codomain, biperm1_domain + algorithm.left_matricize_style, op1, a1, biperm1_codomain, biperm1_domain ) a2_mat = matricizeopperm( - algorithm.right_fusion_style, op2, a2, biperm2_codomain, biperm2_domain + algorithm.right_matricize_style, op2, a2, biperm2_codomain, biperm2_domain ) - output_style = algorithm.output_fusion_style + output_style = algorithm.output_matricize_style if iszero(β) && !matricizepermaliases(output_style, invperm_codomain, invperm_domain) # `β` is a strong zero and matricizing `a_dest` would only build a detached copy that # `mul!` immediately overwrites, so skip that gather: let the matmul allocate its matrix diff --git a/src/contract/contractalgorithm.jl b/src/contract/contractalgorithm.jl index bf7e0d8d..02fb0183 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -4,12 +4,12 @@ ContractAlgorithm(algorithm::ContractAlgorithm) = algorithm struct DefaultContractAlgorithm <: ContractAlgorithm end struct Matricize{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm - left_fusion_style::LeftStyle - right_fusion_style::RightStyle - output_fusion_style::OutputStyle + left_matricize_style::LeftStyle + right_matricize_style::RightStyle + output_matricize_style::OutputStyle end -Matricize(fusion_style) = Matricize(fusion_style, fusion_style, fusion_style) -Matricize() = Matricize(ReshapeFusion()) +Matricize(matricize_style) = Matricize(matricize_style, matricize_style, matricize_style) +Matricize() = Matricize(ReshapeMatricize()) """ TensorOperationsAlgorithm(; backend = nothing, allocator = nothing) @@ -36,5 +36,5 @@ 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(FusionStyle(FusionStyle(A1), FusionStyle(A2))) + return Matricize(MatricizeStyle(MatricizeStyle(A1), MatricizeStyle(A2))) end diff --git a/src/diagonal.jl b/src/diagonal.jl index 48d03bdc..23296a49 100644 --- a/src/diagonal.jl +++ b/src/diagonal.jl @@ -1,6 +1,6 @@ using LinearAlgebra: Diagonal -# `Diagonal` participates in the `ReshapeFusion` interface like a dense matrix (it fuses with +# `Diagonal` participates in the `ReshapeMatricize` interface like a dense matrix (it fuses with # the same row/column reshape order), but its structure is preserved wherever the result of # an operation is still diagonal. These methods hook the lowest-level primitives, so the # convenience wrappers built on them (`bipermutedims`, `permutedimsadd!`, `add!`, and the @@ -44,9 +44,9 @@ end # A `Diagonal` is already a matrix; the `(1 codomain, 1 domain)` matricization is the identity # reshape, so return it directly (maybe-alias, matching `matricize`'s general contract). -matricize(::ReshapeFusion, a::Diagonal, ::Val{1}) = a +matricize(::ReshapeMatricize, a::Diagonal, ::Val{1}) = a function unmatricize( - ::ReshapeFusion, m::Diagonal, + ::ReshapeMatricize, m::Diagonal, ::Tuple{<:AbstractUnitRange}, ::Tuple{<:AbstractUnitRange} ) return m diff --git a/src/directsum.jl b/src/directsum.jl index 9a51eec7..85d257cd 100644 --- a/src/directsum.jl +++ b/src/directsum.jl @@ -1,3 +1,3 @@ # `directsum` is a plain concatenation for now, kept as its own entry point so a fusing/rotating -# variant can later be selected by style, the way `matricize` takes a `FusionStyle`. +# variant can later be selected by style, the way `matricize` takes a `MatricizeStyle`. directsum(dims, as...) = concatenate(dims, as...) diff --git a/src/factorizations.jl b/src/factorizations.jl index 85c2329f..7a6f9a64 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -15,7 +15,7 @@ for f in ( :left_polar, :right_polar, :left_orth, :right_orth, ) @eval begin - function $f(style::FusionStyle, A, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) X, Y = MatrixAlgebraKit.$f(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) @@ -23,7 +23,7 @@ for f in ( unmatricize(style, Y, (axes(Y, 1),), axes_domain) end function $f(A, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(A), A, ndims_codomain; kwargs...) + return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) end end end @@ -38,7 +38,7 @@ for f in ( ) @eval begin function $f( - style::FusionStyle, A, + style::MatricizeStyle, A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs... ) @@ -55,7 +55,7 @@ for f in ( end function $f( - style::FusionStyle, A, + style::MatricizeStyle, A, labels_A, labels_codomain, labels_domain; kwargs... ) perm_codomain, perm_domain = @@ -96,11 +96,11 @@ julia> TensorAlgebra.tr(A, (:i, :j, :k, :l), (:i, :k), (:j, :l)) ≈ true ``` """ -function tr(style::FusionStyle, A, ndims_codomain::Val) +function tr(style::MatricizeStyle, A, ndims_codomain::Val) return LinearAlgebra.tr(matricize(style, A, ndims_codomain)) end function tr(A, ndims_codomain::Val) - return tr(FusionStyle(A), A, ndims_codomain) + return tr(MatricizeStyle(A), A, ndims_codomain) end function tr(A, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}) A_perm = bipermutedims(A, perm_codomain, perm_domain) @@ -256,7 +256,7 @@ right_orth # rank × rank spectrum, and `Vᴴ` carries a leading rank axis plus the domain axes. for f in (:svd_compact, :svd_full) @eval begin - function $f(style::FusionStyle, A, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) U, S, Vᴴ = MatrixAlgebraKit.$f(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) @@ -265,7 +265,7 @@ for f in (:svd_compact, :svd_full) unmatricize(style, Vᴴ, (axes(Vᴴ, 1),), axes_domain) end function $f(A, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(A), A, ndims_codomain; kwargs...) + return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) end end end @@ -273,7 +273,7 @@ end # `svd_trunc` matches the three-output SVD but additionally surfaces the truncation error # `ϵ` (the 2-norm of the discarded singular values, computed by MatrixAlgebraKit without # catastrophic cancellation), so it is spelled out here rather than sharing the loop above. -function svd_trunc(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function svd_trunc(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) U, S, Vᴴ, ϵ = MatrixAlgebraKit.svd_trunc(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) @@ -283,21 +283,21 @@ function svd_trunc(style::FusionStyle, A, ndims_codomain::Val; kwargs...) ϵ end function svd_trunc(A, ndims_codomain::Val; kwargs...) - return svd_trunc(FusionStyle(A), A, ndims_codomain; kwargs...) + return svd_trunc(MatricizeStyle(A), A, ndims_codomain; kwargs...) end # Eigendecomposition: `D` is the rank × rank spectrum, left as a matrix, while `V` carries # the codomain axes plus a trailing rank axis. for f in (:eigh_full, :eig_full, :eigh_trunc, :eig_trunc) @eval begin - function $f(style::FusionStyle, A, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) D, V = MatrixAlgebraKit.$f(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return D, unmatricize(style, V, axes_codomain, (conj(axes(V, ndims(V))),)) end function $f(A, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(A), A, ndims_codomain; kwargs...) + return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) end end end @@ -305,12 +305,12 @@ end # Spectrum-only factorizations returning a vector of singular values / eigenvalues. for f in (:svd_vals, :eigh_vals, :eig_vals) @eval begin - function $f(style::FusionStyle, A, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) return MatrixAlgebraKit.$f(A_mat; kwargs...) end function $f(A, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(A), A, ndims_codomain; kwargs...) + return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) end end end @@ -487,17 +487,17 @@ The output satisfies `N' * A ≈ 0` and `N' * N ≈ I`. """ left_null -function left_null!!(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function left_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) N = MatrixAlgebraKit.left_null!(A_mat; kwargs...) axes_codomain = first(bipartition(axes(A), ndims_codomain)) return unmatricize(style, N, axes_codomain, (conj(axes(N, ndims(N))),)) end function left_null!!(A, ndims_codomain::Val; kwargs...) - return left_null!!(FusionStyle(A), A, ndims_codomain; kwargs...) + return left_null!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end -function left_null(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function left_null(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) return left_null!!(style, copy(A), ndims_codomain; kwargs...) end function left_null(A, ndims_codomain::Val; kwargs...) @@ -524,17 +524,17 @@ The output satisfies `A * Nᴴ' ≈ 0` and `Nᴴ * Nᴴ' ≈ I`. """ right_null -function right_null!!(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function right_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) Nᴴ = MatrixAlgebraKit.right_null!(A_mat; kwargs...) _, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, Nᴴ, (axes(Nᴴ, 1),), axes_domain) end function right_null!!(A, ndims_codomain::Val; kwargs...) - return right_null!!(FusionStyle(A), A, ndims_codomain; kwargs...) + return right_null!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end -function right_null(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function right_null(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) return right_null!!(style, copy(A), ndims_codomain; kwargs...) end function right_null(A, ndims_codomain::Val; kwargs...) @@ -579,7 +579,7 @@ See also [`gram_eigh_full_with_pinv`](@ref) and gram_eigh_full function gram_eigh_full!!( - style::FusionStyle, A, ndims_codomain::Val; kwargs... + style::MatricizeStyle, A, ndims_codomain::Val; kwargs... ) A_mat = matricize(style, A, ndims_codomain) X = MatrixAlgebra.gram_eigh_full!!(A_mat; kwargs...) @@ -587,11 +587,11 @@ function gram_eigh_full!!( 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!!(FusionStyle(A), A, ndims_codomain; kwargs...) + return gram_eigh_full!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end function gram_eigh_full( - style::FusionStyle, A, ndims_codomain::Val; kwargs... + style::MatricizeStyle, A, ndims_codomain::Val; kwargs... ) return gram_eigh_full!!(style, copy(A), ndims_codomain; kwargs...) end @@ -640,7 +640,7 @@ See also [`MatrixAlgebra.gram_eigh_full_with_pinv`](@ref). gram_eigh_full_with_pinv function gram_eigh_full_with_pinv!!( - style::FusionStyle, A, ndims_codomain::Val; kwargs... + 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...) @@ -649,11 +649,11 @@ function gram_eigh_full_with_pinv!!( 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!!(FusionStyle(A), A, ndims_codomain; kwargs...) + return gram_eigh_full_with_pinv!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end function gram_eigh_full_with_pinv( - style::FusionStyle, A, ndims_codomain::Val; kwargs... + style::MatricizeStyle, A, ndims_codomain::Val; kwargs... ) return gram_eigh_full_with_pinv!!(style, copy(A), ndims_codomain; kwargs...) end @@ -709,14 +709,14 @@ invsqrth_safe for f in (:sqrth_safe, :invsqrth_safe) @eval begin - function $f(style::FusionStyle, A, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) P_mat = MatrixAlgebra.$f(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, P_mat, axes_codomain, axes_domain) end function $f(A, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(A), A, ndims_codomain; kwargs...) + return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) end end end @@ -734,14 +734,14 @@ See also `MatrixAlgebraKit.project_hermitian`. """ project_hermitian -function project_hermitian(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function project_hermitian(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) H_mat = MatrixAlgebraKit.project_hermitian(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, H_mat, axes_codomain, axes_domain) end function project_hermitian(A, ndims_codomain::Val; kwargs...) - return project_hermitian(FusionStyle(A), A, ndims_codomain; kwargs...) + return project_hermitian(MatricizeStyle(A), A, ndims_codomain; kwargs...) end """ @@ -764,7 +764,7 @@ See also [`MatrixAlgebra.sqrth_invsqrth_safe`](@ref). """ sqrth_invsqrth_safe -function sqrth_invsqrth_safe(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function sqrth_invsqrth_safe(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) P_mat, Pinv_mat = MatrixAlgebra.sqrth_invsqrth_safe(A_mat; kwargs...) axes_codomain, axes_domain = bipartition_axes(axes(A), ndims_codomain) @@ -772,7 +772,7 @@ function sqrth_invsqrth_safe(style::FusionStyle, A, ndims_codomain::Val; kwargs. unmatricize(style, Pinv_mat, axes_codomain, axes_domain) end function sqrth_invsqrth_safe(A, ndims_codomain::Val; kwargs...) - return sqrth_invsqrth_safe(FusionStyle(A), A, ndims_codomain; kwargs...) + return sqrth_invsqrth_safe(MatricizeStyle(A), A, ndims_codomain; kwargs...) end """ @@ -808,30 +808,30 @@ true """ one -function one!!(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function one!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) MatrixAlgebraKit.one!(A_mat) codomain_axes, domain_axes = bipartition_axes(axes(A), ndims_codomain) return unmatricize(style, A_mat, codomain_axes, domain_axes) end function one!!(A, ndims_codomain::Val; kwargs...) - return one!!(FusionStyle(A), A, ndims_codomain; kwargs...) + return one!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end # In-place identity fill: writes the identity into `A` and returns it. Matricizes `A`, fills the # fused matrix with the identity, and — when the matricized form is a detached copy (a graded # gather) rather than a view aliasing `A` (a dense reshape) — scatters it back with `unmatricize!`. -function one!(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function one!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) A_mat = matricize(style, A, ndims_codomain) MatrixAlgebraKit.one!(A_mat) Base.mightalias(A_mat, A) && return A return unmatricize!(A, A_mat, ndims_codomain) end function one!(A, ndims_codomain::Val; kwargs...) - return one!(FusionStyle(A), A, ndims_codomain; kwargs...) + return one!(MatricizeStyle(A), A, ndims_codomain; kwargs...) end -function one(style::FusionStyle, A, ndims_codomain::Val; kwargs...) +function one(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) return one!!(style, copy(A), ndims_codomain; kwargs...) end function one(A, ndims_codomain::Val; kwargs...) diff --git a/src/matricize.jl b/src/matricize.jl index 9113b7fc..a86bf8fc 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -1,13 +1,13 @@ using EllipsisNotation: Ellipsis using LinearAlgebra: Diagonal -# ===================================== FusionStyle ====================================== -abstract type FusionStyle end +# ===================================== MatricizeStyle ====================================== +abstract type MatricizeStyle end -FusionStyle(x) = FusionStyle(typeof(x)) -FusionStyle(T::Type) = throw(MethodError(FusionStyle, (T,))) -FusionStyle(style1::Style, style2::Style) where {Style <: FusionStyle} = Style() -FusionStyle(style1::FusionStyle, style2::FusionStyle) = ReshapeFusion() +MatricizeStyle(x) = MatricizeStyle(typeof(x)) +MatricizeStyle(T::Type) = throw(MethodError(MatricizeStyle, (T,))) +MatricizeStyle(style1::Style, style2::Style) where {Style <: MatricizeStyle} = Style() +MatricizeStyle(style1::MatricizeStyle, style2::MatricizeStyle) = ReshapeMatricize() # ======================================= misc ======================================== @@ -72,28 +72,28 @@ end # matrix factorizations assume copy # maybe: copy=false kwarg -# This is the primary function that should be overloaded for new fusion styles. +# This is the primary function that should be overloaded for new matricize styles. # This assumes the permutation was already performed. function matricize( - style::FusionStyle, a, ndims_codomain::Val + style::MatricizeStyle, a, ndims_codomain::Val ) return throw(MethodError(matricize, (style, a, ndims_codomain))) end function matricize(a, ndims_codomain::Val) - return matricize(FusionStyle(a), a, ndims_codomain) + return matricize(MatricizeStyle(a), a, ndims_codomain) end function matricizeperm( a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) - return matricizeperm(FusionStyle(a), a, perm_codomain, perm_domain) + return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) end # Thin wrapper around `matricizeopperm` with identity op — the actual matricization logic -# (and the fusion-style overload point for folding ops into matricization) lives in +# (and the matricize-style overload point for folding ops into matricization) lives in # `matricizeopperm`. function matricizeperm( - style::FusionStyle, a, + style::MatricizeStyle, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) return matricizeopperm(style, identity, a, perm_codomain, perm_domain) @@ -124,10 +124,10 @@ function to_permblocks( end function matricizeperm(a, perm_codomain, perm_domain) - return matricizeperm(FusionStyle(a), a, perm_codomain, perm_domain) + return matricizeperm(MatricizeStyle(a), a, perm_codomain, perm_domain) end function matricizeperm( - style::FusionStyle, a, perm_codomain, perm_domain + style::MatricizeStyle, a, perm_codomain, perm_domain ) return matricizeperm(style, a, to_permblocks(a, (perm_codomain, perm_domain))...) end @@ -141,14 +141,14 @@ Matricize `a` with element-wise operation `op` folded in. Returns a matrix repre `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 fusion style and array type. The caller should treat the result +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(FusionStyle(a), op, a, perm_codomain, perm_domain) + return matricizeopperm(MatricizeStyle(a), op, a, perm_codomain, perm_domain) end function matricizeopperm( - style::FusionStyle, op, a, perm_codomain, perm_domain + style::MatricizeStyle, op, a, perm_codomain, perm_domain ) return matricizeopperm(style, op, a, to_permblocks(a, (perm_codomain, perm_domain))...) end @@ -162,17 +162,17 @@ end # array realizes as a `transpose` of a `reshape` (a view gemm # reads via BLAS' transpose flag). # PermuteMatricizeKind — the groups interleave storage, so a permuted copy is required. -# Pure: depends only on the index pattern, not on `a`'s data. Dispatched on `FusionStyle`. +# Pure: depends only on the index pattern, not on `a`'s data. Dispatched on `MatricizeStyle`. # The generic classifier only recognizes the always-safe `ReshapeMatricizeKind` (skipping a # no-op permute is valid for any style); `TransposeMatricizeKind` is opt-in for styles whose -# `matricize` composes with a lazy `transpose`, currently only `ReshapeFusion`. +# `matricize` composes with a lazy `transpose`, currently only `ReshapeMatricize`. @enum MatricizeKind ReshapeMatricizeKind TransposeMatricizeKind PermuteMatricizeKind # Whether `perm` is the identity permutation `(1, …, n)`. isidentityperm(perm::Tuple{Vararg{Int}}) = perm == ntuple(identity, length(perm)) function matricizekind( - ::FusionStyle, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + ::MatricizeStyle, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) # Already in storage order: the permute is a no-op, so `matricize` can run directly. isidentityperm((perm_codomain..., perm_domain...)) && return ReshapeMatricizeKind @@ -184,7 +184,7 @@ end # `matricize` is itself a view (a dense reshape) can alias, and only when the bipermutation needs # no permuted copy. Defaults to `false`: a style that gathers into new storage, such as a graded # array, never aliases its input. -matricizepermaliases(::FusionStyle, perm_codomain, perm_domain) = false +matricizepermaliases(::MatricizeStyle, perm_codomain, perm_domain) = false # Skip the permuted copy when the classifier says it is unnecessary. `ReshapeMatricizeKind` # calls `matricize` directly on `a` (a view for dense, a gather without the extra permute @@ -192,7 +192,7 @@ matricizepermaliases(::FusionStyle, perm_codomain, perm_domain) = false # fast paths require `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::FusionStyle, op, a, + style::MatricizeStyle, op, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) ndims(a) == length(perm_codomain) + length(perm_domain) || @@ -211,16 +211,16 @@ end # ==================================== unmatricize ======================================= # Split form: `codomain_axes` and `domain_axes` are the destination axes for the codomain and # domain groups, given codomain-facing (un-dualized), the same convention as `similar_map`. A -# fusion 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 fusion styles. +# 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 # disambiguate axis tuples from permutation tuples regardless of how unconstrained `m` and the # axes are. -function unmatricize(style::FusionStyle, m, codomain_axes, domain_axes) +function unmatricize(style::MatricizeStyle, m, codomain_axes, domain_axes) return throw(MethodError(unmatricize, (style, m, codomain_axes, domain_axes))) end function unmatricize(m, codomain_axes, domain_axes) - return unmatricize(FusionStyle(m), m, codomain_axes, domain_axes) + return unmatricize(MatricizeStyle(m), m, codomain_axes, domain_axes) end # Split `axes` into its codomain and domain groups like `bipartition`, but present the domain @@ -238,10 +238,16 @@ function unmatricizeperm( m, axes_dest, invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} ) - return unmatricizeperm(FusionStyle(m), m, axes_dest, invperm_codomain, invperm_domain) + return unmatricizeperm( + MatricizeStyle(m), + m, + axes_dest, + invperm_codomain, + invperm_domain + ) end function unmatricizeperm( - style::FusionStyle, m, axes_dest, + style::MatricizeStyle, m, axes_dest, invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} ) invbiperm = BiTuple(invperm_codomain, invperm_domain) @@ -257,10 +263,10 @@ function unmatricizeperm!( a_dest, m, invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} ) - return unmatricizeperm!(FusionStyle(m), a_dest, m, invperm_codomain, invperm_domain) + return unmatricizeperm!(MatricizeStyle(m), a_dest, m, invperm_codomain, invperm_domain) end function unmatricizeperm!( - style::FusionStyle, a_dest, m, + style::MatricizeStyle, a_dest, m, invperm_codomain::Tuple{Vararg{Int}}, invperm_domain::Tuple{Vararg{Int}} ) invbiperm = BiTuple(invperm_codomain, invperm_domain) @@ -287,10 +293,10 @@ function unmatricize!(a_dest, m, ndims_codomain::Val) ) end -# Defaults to ReshapeFusion, a simple reshape -struct ReshapeFusion <: FusionStyle end -FusionStyle(::Type{<:AbstractArray}) = ReshapeFusion() -function matricize(::ReshapeFusion, a, ndims_codomain::Val) +# Defaults to ReshapeMatricize, a simple reshape +struct ReshapeMatricize <: MatricizeStyle end +MatricizeStyle(::Type{<:AbstractArray}) = ReshapeMatricize() +function matricize(::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) @@ -300,7 +306,7 @@ end # reshape (a view), so it opts into `TransposeMatricizeKind` on top of the generic # reshape/permute classification. function matricizekind( - ::ReshapeFusion, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} + ::ReshapeMatricize, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}} ) isidentityperm((perm_codomain..., perm_domain...)) && return ReshapeMatricizeKind isidentityperm((perm_domain..., perm_codomain...)) && return TransposeMatricizeKind @@ -308,11 +314,11 @@ function matricizekind( end # A dense reshape/transpose is a view of `a`; only a permuted copy detaches. So the matricized # output aliases `a` for every kind except `PermuteMatricizeKind`. -function matricizepermaliases(style::ReshapeFusion, perm_codomain, perm_domain) +function matricizepermaliases(style::ReshapeMatricize, perm_codomain, perm_domain) return matricizekind(style, perm_codomain, perm_domain) != PermuteMatricizeKind end # A dense reshape ignores the codomain/domain split: it just reshapes to the concatenated axes. # `conj` re-dualizes the codomain-facing `domain_axes` into stored form, a no-op on a dense axis. -function unmatricize(style::ReshapeFusion, m, codomain_axes, domain_axes) +function unmatricize(style::ReshapeMatricize, m, codomain_axes, domain_axes) return reshape(m, (codomain_axes..., conj.(domain_axes)...)) end diff --git a/src/matrixfunctions.jl b/src/matrixfunctions.jl index 1734a9e2..85240c08 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -33,18 +33,18 @@ const MATRIX_FUNCTIONS = [ for f in MATRIX_FUNCTIONS @eval begin - function $f(style::FusionStyle, a, ndims_codomain::Val; kwargs...) + function $f(style::MatricizeStyle, a, ndims_codomain::Val; kwargs...) a_mat = matricize(style, a, ndims_codomain) fa_mat = Base.$f(a_mat; kwargs...) codomain_axes, domain_axes = bipartition_axes(axes(a), ndims_codomain) return unmatricize(style, fa_mat, codomain_axes, domain_axes) end function $f(a, ndims_codomain::Val; kwargs...) - return $f(FusionStyle(a), a, ndims_codomain; kwargs...) + return $f(MatricizeStyle(a), a, ndims_codomain; kwargs...) end function $f( - style::FusionStyle, a, + style::MatricizeStyle, a, perm_codomain::Tuple{Vararg{Int}}, perm_domain::Tuple{Vararg{Int}}; kwargs... ) @@ -61,7 +61,7 @@ for f in MATRIX_FUNCTIONS end function $f( - style::FusionStyle, a, + style::MatricizeStyle, a, labels_a, labels_codomain, labels_domain; kwargs... ) perm_codomain, perm_domain = diff --git a/test/Project.toml b/test/Project.toml index 6e25853d..aed429d1 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -37,7 +37,7 @@ Random = "1.10" SafeTestsets = "0.1" StableRNGs = "1.0.2" Suppressor = "0.2" -TensorAlgebra = "0.18" +TensorAlgebra = "0.19" TensorKit = "0.17" TensorOperations = "5.1.4" Test = "1.10" diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index 966e67ec..29795c6d 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -41,13 +41,13 @@ using Test: @test, @testset end @testset "matricize(1, 1) is the identity reshape" begin - m = TensorAlgebra.matricize(TensorAlgebra.ReshapeFusion(), d, Val(1)) + m = TensorAlgebra.matricize(TensorAlgebra.ReshapeMatricize(), d, Val(1)) @test m === d end @testset "unmatricize round-trips a Diagonal" begin ax = axes(d, 1) - back = TensorAlgebra.unmatricize(TensorAlgebra.ReshapeFusion(), d, (ax,), (ax,)) + back = TensorAlgebra.unmatricize(TensorAlgebra.ReshapeMatricize(), d, (ax,), (ax,)) @test back === d end diff --git a/test/test_fusionstyle.jl b/test/test_fusionstyle.jl deleted file mode 100644 index a37a8128..00000000 --- a/test/test_fusionstyle.jl +++ /dev/null @@ -1,28 +0,0 @@ -using TensorAlgebra: TensorAlgebra as TA, FusionStyle, Matricize, ReshapeFusion -using Test: @test, @testset - -module FusionStyleTestUtils - using TensorAlgebra: TensorAlgebra as TA - struct MyArray{T, N, A <: AbstractArray{T, N}} <: AbstractArray{T, N} - parent::A - end - struct MyArrayFusion <: TA.FusionStyle end - TA.FusionStyle(::Type{<:MyArray}) = MyArrayFusion() -end -using .FusionStyleTestUtils: MyArray, MyArrayFusion - -@testset "FusionStyle" begin - a1 = randn(2, 2) - a2 = MyArray(randn(2, 2)) - @test FusionStyle(a1) ≡ ReshapeFusion() - @test FusionStyle(a2) ≡ MyArrayFusion() - @test FusionStyle(typeof(a1)) ≡ ReshapeFusion() - @test FusionStyle(ReshapeFusion(), ReshapeFusion()) ≡ ReshapeFusion() - @test FusionStyle(MyArrayFusion(), MyArrayFusion()) ≡ MyArrayFusion() - @test FusionStyle(MyArrayFusion(), ReshapeFusion()) ≡ ReshapeFusion() - @test FusionStyle(ReshapeFusion(), MyArrayFusion()) ≡ ReshapeFusion() - @test TA.default_contract_algorithm(typeof(a1), typeof(a1)) ≡ Matricize(ReshapeFusion()) - @test TA.default_contract_algorithm(typeof(a1), typeof(a2)) ≡ Matricize(ReshapeFusion()) - @test TA.default_contract_algorithm(typeof(a2), typeof(a1)) ≡ Matricize(ReshapeFusion()) - @test TA.default_contract_algorithm(typeof(a2), typeof(a2)) ≡ Matricize(MyArrayFusion()) -end diff --git a/test/test_matricize.jl b/test/test_matricize.jl index 7294ef09..bb9139a1 100644 --- a/test/test_matricize.jl +++ b/test/test_matricize.jl @@ -1,11 +1,11 @@ using LinearAlgebra: Transpose using StableRNGs: StableRNG -using TensorAlgebra: TensorAlgebra, PermuteMatricizeKind, ReshapeFusion, +using TensorAlgebra: TensorAlgebra, PermuteMatricizeKind, ReshapeMatricize, ReshapeMatricizeKind, TransposeMatricizeKind, matricizeopperm, matricizeperm using Test: @test, @testset -# A non-`ReshapeFusion` style, to check the always-safe generic fallback. -struct DummyFusion <: TensorAlgebra.FusionStyle end +# A non-`ReshapeMatricize` style, to check the always-safe generic fallback. +struct DummyMatricize <: TensorAlgebra.MatricizeStyle end # Ground-truth matricization: permute into `(codomain..., domain...)` order, then reshape. function matricize_ref(a, perm_codomain, perm_domain) @@ -16,7 +16,7 @@ function matricize_ref(a, perm_codomain, perm_domain) end @testset "matricizekind classifier" begin - style = ReshapeFusion() + style = ReshapeMatricize() # Already in storage order → plain reshape view. @test TensorAlgebra.matricizekind(style, (1,), (2, 3)) == ReshapeMatricizeKind @test TensorAlgebra.matricizekind(style, (1, 2), (3,)) == ReshapeMatricizeKind @@ -29,12 +29,16 @@ end @test TensorAlgebra.matricizekind(style, (3, 1), (2,)) == PermuteMatricizeKind @test TensorAlgebra.matricizekind(style, (2,), (1, 3)) == PermuteMatricizeKind @test TensorAlgebra.matricizekind(style, (1, 3), (2,)) == PermuteMatricizeKind - # Generic fusion styles recognize the always-safe reshape (no-op permute) but never + # Generic matricize styles recognize the always-safe reshape (no-op permute) but never # claim a transpose (which only styles with a lazy `transpose` can realize). - @test TensorAlgebra.matricizekind(DummyFusion(), (1,), (2, 3)) == ReshapeMatricizeKind - @test TensorAlgebra.matricizekind(DummyFusion(), (1, 2, 3), ()) == ReshapeMatricizeKind - @test TensorAlgebra.matricizekind(DummyFusion(), (2, 3), (1,)) == PermuteMatricizeKind - @test TensorAlgebra.matricizekind(DummyFusion(), (3, 1), (2,)) == PermuteMatricizeKind + @test TensorAlgebra.matricizekind(DummyMatricize(), (1,), (2, 3)) == + ReshapeMatricizeKind + @test TensorAlgebra.matricizekind(DummyMatricize(), (1, 2, 3), ()) == + ReshapeMatricizeKind + @test TensorAlgebra.matricizekind(DummyMatricize(), (2, 3), (1,)) == + PermuteMatricizeKind + @test TensorAlgebra.matricizekind(DummyMatricize(), (3, 1), (2,)) == + PermuteMatricizeKind end @testset "maybe-view matricizeopperm (eltype=$elt)" for elt in (Float64, ComplexF64) diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl new file mode 100644 index 00000000..9ff7dd37 --- /dev/null +++ b/test/test_matricizestyle.jl @@ -0,0 +1,32 @@ +using TensorAlgebra: TensorAlgebra as TA, Matricize, MatricizeStyle, ReshapeMatricize +using Test: @test, @testset + +module MatricizeStyleTestUtils + using TensorAlgebra: TensorAlgebra as TA + struct MyArray{T, N, A <: AbstractArray{T, N}} <: AbstractArray{T, N} + parent::A + end + struct MyArrayMatricize <: TA.MatricizeStyle end + TA.MatricizeStyle(::Type{<:MyArray}) = MyArrayMatricize() +end +using .MatricizeStyleTestUtils: MyArray, MyArrayMatricize + +@testset "MatricizeStyle" begin + a1 = randn(2, 2) + a2 = MyArray(randn(2, 2)) + @test MatricizeStyle(a1) ≡ ReshapeMatricize() + @test MatricizeStyle(a2) ≡ MyArrayMatricize() + @test MatricizeStyle(typeof(a1)) ≡ ReshapeMatricize() + @test MatricizeStyle(ReshapeMatricize(), ReshapeMatricize()) ≡ ReshapeMatricize() + @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()) +end