diff --git a/Project.toml b/Project.toml index 29a6dcc6..419a01b6 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.21.1" +version = "0.22.0" authors = ["ITensor developers and contributors"] [workspace] diff --git a/docs/Project.toml b/docs/Project.toml index d9d7449e..020ddb86 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.21" +TensorAlgebra = "0.22" diff --git a/examples/Project.toml b/examples/Project.toml index 194d5b6d..fde811e6 100644 --- a/examples/Project.toml +++ b/examples/Project.toml @@ -5,4 +5,4 @@ TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" path = ".." [compat] -TensorAlgebra = "0.21" +TensorAlgebra = "0.22" diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index 608f8d4c..e229daa4 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -241,14 +241,12 @@ 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 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 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, + ::typeof(TensorAlgebra.matricizeop), op, t::AbstractTensorMap, perm_codomain, perm_domain ) return op === identity && @@ -256,7 +254,7 @@ function TensorAlgebra.is_output_view( length(perm_codomain) == numout(t) end function TensorAlgebra.matricizeopview( - ::TensorKitMatricize, op, t::AbstractTensorMap, perm_codomain, perm_domain + op, t::AbstractTensorMap, perm_codomain, perm_domain ) return t end @@ -265,7 +263,7 @@ end # `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, + ::typeof(TensorAlgebra.matricizeop), op, t::AbstractTensorMap, perm_codomain, perm_domain ) return TensorAlgebra.allocate_output( @@ -273,7 +271,7 @@ function TensorAlgebra.allocate_output( ) end function TensorAlgebra.matricizeop!( - dest::AbstractTensorMap, ::TensorKitMatricize, op, + dest::AbstractTensorMap, op, t::AbstractTensorMap, perm_codomain, perm_domain ) return TensorAlgebra.bipermutedimsopadd!( @@ -285,9 +283,7 @@ end # 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 - ) +function TensorAlgebra.unmatricize(m::AbstractTensorMap, axes_codomain, axes_domain) S = spacetype(m) dest = ProductSpace{S}(axes_codomain...) ← ProductSpace{S}(axes_domain...) space(m) == dest || diff --git a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl index 0bcb5134..0dd50baa 100644 --- a/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl +++ b/ext/TensorAlgebraTensorOperationsExt/TensorAlgebraTensorOperationsExt.jl @@ -30,6 +30,12 @@ function TA.contractpermopadd!( op2, a2, perm2_codomain, perm2_domain, α::Number, β::Number ) + TA.check_input( + TA.contract!, + a_dest, perm_dest_codomain, perm_dest_domain, + a1, perm1_codomain, perm1_domain, + a2, 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)) diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 92887326..8902fb80 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -5,7 +5,7 @@ export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public AbstractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermalign, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, has_bipartition, 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" + "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, has_bipartition, 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, 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, unmatricizeadd!, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/src/contract/contract.jl b/src/contract/contract.jl index 1cf720ed..7483251b 100644 --- a/src/contract/contract.jl +++ b/src/contract/contract.jl @@ -211,12 +211,8 @@ function contractpermopadd!( α::Number, β::Number; alg = nothing ) - check_input( - contract!, - a_dest, perm_dest_codomain, perm_dest_domain, - a1, perm1_codomain, perm1_domain, - a2, perm2_codomain, perm2_domain - ) + # Input validation belongs to the algorithm's kernel below, which is a public rung callable + # directly, so this entry only selects and dispatches. algorithm = select_algorithm(contract!, alg, (a_dest, a1, a2)) return contractpermopadd!( algorithm, @@ -253,12 +249,12 @@ end # 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. +# matching `check_input`. A backend can therefore choose on the destination as well as on the +# operands. 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))) + ::typeof(contract!), ::Type{<:Tuple{AbstractArray, AbstractArray, AbstractArray}} + ) + return MatricizeContract() end function select_algorithm(::typeof(contract!), ::DefaultContractAlgorithm, args::Tuple) return default_algorithm(contract!, args) diff --git a/src/contract/contract_matricize.jl b/src/contract/contract_matricize.jl index 7b4e66e3..67973a38 100644 --- a/src/contract/contract_matricize.jl +++ b/src/contract/contract_matricize.jl @@ -1,7 +1,11 @@ using LinearAlgebra: mul! +# The matricized kernel for arrays whose matricization is a plain fold: matricize both operands, +# multiply, and write the product into the destination's matricization. An array family whose +# contraction needs more than the fold (a fermionic twist, a block-sparse product) owns its own +# `ContractAlgorithm` and `contractpermopadd!` method rather than hooking into this one. function contractpermopadd!( - algorithm::MatricizeContract, + ::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,37 +20,19 @@ function contractpermopadd!( a1, biperm1_codomain, biperm1_domain, a2, biperm2_codomain, biperm2_domain ) - a1_mat = matricizeop( - algorithm.left_matricize_style, op1, a1, biperm1_codomain, biperm1_domain - ) - a2_mat = matricizeop( - algorithm.right_matricize_style, op2, a2, biperm2_codomain, biperm2_domain - ) - output_style = algorithm.output_matricize_style - if is_output_view( - matricizeop, output_style, identity, a_dest, invperm_codomain, invperm_domain - ) + a1_mat = matricizeop(op1, a1, biperm1_codomain, biperm1_domain) + a2_mat = matricizeop(op2, a2, biperm2_codomain, biperm2_domain) + if is_output_view(matricizeop, identity, a_dest, invperm_codomain, invperm_domain) # The matricization shares `a_dest`'s memory, so the matmul is the whole operation. - a_dest_mat = matricizeopview( - output_style, identity, a_dest, invperm_codomain, invperm_domain - ) + a_dest_mat = matricizeopview(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 - # allocate its matrix result and scatter it into `a_dest`. Every coupled-sector block - # is materialized (the matmul zeros the ones it does not reach), so the scatter - # overwrites `a_dest` in full. - a_dest_mat = a1_mat * a2_mat - isone(α) || scale!(a_dest_mat, α) - 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 = matricizeopcopy( - output_style, identity, a_dest, invperm_codomain, invperm_domain - ) - mul!(a_dest_mat, a1_mat, a2_mat, α, β) - unmatricize!(output_style, a_dest, a_dest_mat, invperm_codomain, invperm_domain) + # Let the matmul allocate its matrix result and scatter it into `a_dest` with `α` and `β` + # folded into the one permuted pass, so `a_dest` is never gathered. Every coupled-sector + # block is materialized (the matmul zeros the ones it does not reach), so the scatter + # reaches `a_dest` in full. + a_dest_mat = a1_mat * a2_mat + unmatricizeadd!(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 e7ffc8a7..73871d38 100644 --- a/src/contract/contractalgorithm.jl +++ b/src/contract/contractalgorithm.jl @@ -3,15 +3,7 @@ ContractAlgorithm(algorithm::ContractAlgorithm) = algorithm struct DefaultContractAlgorithm <: ContractAlgorithm end -struct MatricizeContract{LeftStyle, RightStyle, OutputStyle} <: ContractAlgorithm - left_matricize_style::LeftStyle - right_matricize_style::RightStyle - output_matricize_style::OutputStyle -end -function MatricizeContract(matricize_style) - return MatricizeContract(matricize_style, matricize_style, matricize_style) -end -MatricizeContract() = MatricizeContract(ReshapeMatricize()) +struct MatricizeContract <: ContractAlgorithm end """ TensorOperationsContract(; backend = nothing, allocator = nothing) diff --git a/src/diagonal.jl b/src/diagonal.jl index c5e3b4a4..3ce43938 100644 --- a/src/diagonal.jl +++ b/src/diagonal.jl @@ -1,6 +1,6 @@ using LinearAlgebra: Diagonal -# `Diagonal` participates in the `ReshapeMatricize` interface like a dense matrix (it fuses with +# `Diagonal` matricizes through the dense reshape hooks like any other 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 @@ -43,9 +43,15 @@ 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). +# `Diagonal`-specialized consumers downstream). Any other split densifies through the copy path. +function is_output_view( + ::typeof(matricizeop), op, a::Diagonal, perm_codomain::Tuple{Int}, + perm_domain::Tuple{Int} + ) + return op === identity && isidentitybiperm(perm_codomain, perm_domain) +end function matricizeopview( - ::ReshapeMatricize, op, a::Diagonal, perm_codomain::Tuple{Int}, perm_domain::Tuple{Int} + op, a::Diagonal, perm_codomain::Tuple{Int}, perm_domain::Tuple{Int} ) return a end @@ -53,7 +59,7 @@ end # result stays `Diagonal`, so return `m` directly. The generic `check_input(unmatricize, ...)` # validates the axis lengths against `m`'s size. function unmatricize( - ::ReshapeMatricize, m::Diagonal, + m::Diagonal, axes_codomain::Tuple{<:AbstractUnitRange}, axes_domain::Tuple{<:AbstractUnitRange} ) check_input(unmatricize, m, axes_codomain, axes_domain) @@ -63,10 +69,8 @@ end # result is not representable as a `Diagonal`, so densify and reshape like a dense matrix. # `copyto!(similar(m, axes(m)), m)` densifies while preserving `m`'s array backend (a plain # `Array` would force the result onto the CPU). -function unmatricize( - style::ReshapeMatricize, m::Diagonal, axes_codomain::Tuple, axes_domain::Tuple - ) - return unmatricize(style, copyto!(similar(m, axes(m)), m), axes_codomain, axes_domain) +function unmatricize(m::Diagonal, axes_codomain::Tuple, axes_domain::Tuple) + return unmatricize(copyto!(similar(m, axes(m)), m), axes_codomain, axes_domain) end # Contracting two `Diagonal`s to a `{1,1}` destination is the matmul/endomorphism pattern diff --git a/src/directsum.jl b/src/directsum.jl index 85d257cd..b9f38d15 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 `MatricizeStyle`. +# variant can later be dispatched on the array type, the way `matricize` is. directsum(dims, as...) = concatenate(dims, as...) diff --git a/src/factorizations.jl b/src/factorizations.jl index be1dd13b..38ae9251 100644 --- a/src/factorizations.jl +++ b/src/factorizations.jl @@ -9,9 +9,8 @@ using MatrixAlgebraKit: MatrixAlgebraKit # bond is dualized to codomain-facing form (`conj`, a no-op on a dense axis) when it lands on the # domain side of the reconstruction, matching the `unmatricize`/`similar_map` axis convention. -# `unmatricize_factors(f, style, F, axes_codomain, axes_domain)` unfolds the matrix-level -# factors `F` of `f` onto the split axes (in the `unmatricize` convention, domain axes -# un-dualized). +# `unmatricize_factors(f, F, axes_codomain, axes_domain)` unfolds the matrix-level factors `F` +# of `f` onto the split axes (in the `unmatricize` convention, domain axes un-dualized). # # 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 — @@ -27,23 +26,15 @@ for f in ( :left_null, :right_null, :project_hermitian, ) @eval begin - function $f( - style::MatricizeStyle, A, - perm_codomain, perm_domain; - kwargs... - ) + function $f(A, perm_codomain, perm_domain; kwargs...) ndims(A) == length(perm_codomain) + length(perm_domain) || throw(ArgumentError("Invalid bipermutation")) A_mat = - if is_output_view( - matricizeop, style, identity, A, perm_codomain, perm_domain - ) - A_shared = - matricizeopview(style, identity, A, perm_codomain, perm_domain) + if is_output_view(matricizeop, identity, A, perm_codomain, perm_domain) + A_shared = matricizeopview(identity, A, perm_codomain, perm_domain) MatrixAlgebraKit.copy_input(MatrixAlgebraKit.$f, A_shared) else - A_gather = - matricizeopcopy(style, identity, A, perm_codomain, perm_domain) + A_gather = matricizeopcopy(identity, A, perm_codomain, perm_domain) if eltype(A_gather) === float(eltype(A_gather)) A_gather else @@ -55,7 +46,7 @@ for f in ( map(i -> axes(A, i), (perm_codomain..., perm_domain...)), Val(length(perm_codomain)) ) - return unmatricize_factors($f, style, F, axes_codomain, axes_domain) + return unmatricize_factors($f, F, axes_codomain, axes_domain) end end end @@ -66,24 +57,20 @@ for f in ( :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, ) @eval begin - function $f( - style::MatricizeStyle, A, - perm_codomain, perm_domain; - kwargs... - ) - A_mat = matricize(style, A, perm_codomain, perm_domain) + function $f(A, perm_codomain, perm_domain; kwargs...) + A_mat = matricize(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...)), Val(length(perm_codomain)) ) - return unmatricize_factors($f, style, F, axes_codomain, axes_domain) + return unmatricize_factors($f, F, axes_codomain, axes_domain) end end end -# The `Val`, style-inferring, and labels forms of both tiers are thin forwarders into the perm -# form, at the identity bipermutation for the `Val` form. +# The `Val` and labels forms of both tiers are thin forwarders into the perm form, at the +# identity bipermutation for the `Val` form. for f in ( :qr_compact, :qr_full, :lq_compact, :lq_full, :left_polar, :right_polar, :left_orth, :right_orth, @@ -93,28 +80,8 @@ for f in ( :sqrth_safe, :invsqrth_safe, :sqrth_invsqrth_safe, :project_hermitian, ) @eval begin - function $f(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - return $f( - style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...; kwargs... - ) - end function $f(A, ndims_codomain::Val; kwargs...) - return $f(MatricizeStyle(A), A, ndims_codomain; kwargs...) - end - function $f( - A, - perm_codomain, perm_domain; - kwargs... - ) - return $f(MatricizeStyle(A), A, perm_codomain, perm_domain; kwargs...) - end - function $f( - style::MatricizeStyle, A, - labels_A, labels_codomain, labels_domain; kwargs... - ) - perm_codomain, perm_domain = - biperm(Tuple.((labels_A, labels_codomain, labels_domain))...) - return $f(style, A, perm_codomain, perm_domain; kwargs...) + return $f(A, identitybiperm(ndims_codomain, Val(ndims(A)))...; kwargs...) end function $f(A, labels_A, labels_codomain, labels_domain; kwargs...) perm_codomain, perm_domain = @@ -131,13 +98,10 @@ for f in ( :left_polar, :right_polar, :left_orth, :right_orth, ) @eval begin - function unmatricize_factors( - ::typeof($f), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) + function unmatricize_factors(::typeof($f), 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_domain) + return unmatricize(X, axes_codomain, (conj(axes(X, ndims(X))),)), + unmatricize(Y, (axes(Y, 1),), axes_domain) end end end @@ -168,13 +132,8 @@ julia> TensorAlgebra.tr(A, (:i, :j, :k, :l), (:i, :k), (:j, :l)) ≈ true ``` """ -function tr(style::MatricizeStyle, A, ndims_codomain::Val) - 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) + return tr(A, identitybiperm(ndims_codomain, Val(ndims(A)))...) end function tr(A, perm_codomain, perm_domain) return LinearAlgebra.tr(matricize(A, perm_codomain, perm_domain)) @@ -329,14 +288,11 @@ 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 unmatricize_factors( - ::typeof($f), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) + function unmatricize_factors(::typeof($f), F, axes_codomain, axes_domain) U, S, Vᴴ = F - return unmatricize(style, U, axes_codomain, (conj(axes(U, ndims(U))),)), + return unmatricize(U, axes_codomain, (conj(axes(U, ndims(U))),)), S, - unmatricize(style, Vᴴ, (axes(Vᴴ, 1),), axes_domain) + unmatricize(Vᴴ, (axes(Vᴴ, 1),), axes_domain) end end end @@ -344,14 +300,11 @@ 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 unmatricize_factors( - ::typeof(svd_trunc), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) +function unmatricize_factors(::typeof(svd_trunc), F, axes_codomain, axes_domain) U, S, Vᴴ, ϵ = F - return unmatricize(style, U, axes_codomain, (conj(axes(U, ndims(U))),)), + return unmatricize(U, axes_codomain, (conj(axes(U, ndims(U))),)), S, - unmatricize(style, Vᴴ, (axes(Vᴴ, 1),), axes_domain), + unmatricize(Vᴴ, (axes(Vᴴ, 1),), axes_domain), ϵ end @@ -360,13 +313,9 @@ end # unfold); `V` is unmatricized back to the array type, as in `svd_*`. for f in (:eigh_full, :eig_full, :eigh_trunc, :eig_trunc) @eval begin - function unmatricize_factors( - ::typeof($f), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) + function unmatricize_factors(::typeof($f), F, axes_codomain, axes_domain) D, V = F - return D, - unmatricize(style, V, axes_codomain, (conj(axes(V, ndims(V))),)) + return D, unmatricize(V, axes_codomain, (conj(axes(V, ndims(V))),)) end end end @@ -375,10 +324,7 @@ end # nothing to unfold. for f in (:svd_vals, :eigh_vals, :eig_vals) @eval begin - function unmatricize_factors( - ::typeof($f), ::MatricizeStyle, F, - axes_codomain, axes_domain - ) + function unmatricize_factors(::typeof($f), F, axes_codomain, axes_domain) return F end end @@ -556,21 +502,15 @@ The output satisfies `N' * A ≈ 0` and `N' * N ≈ I`. """ left_null -function left_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) +function left_null!!(A, ndims_codomain::Val; kwargs...) + A_mat = matricize(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))),)) -end -function left_null!!(A, ndims_codomain::Val; kwargs...) - return left_null!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) + return unmatricize(N, axes_codomain, (conj(axes(N, ndims(N))),)) end -function unmatricize_factors( - ::typeof(left_null), style::MatricizeStyle, N, - axes_codomain, axes_domain - ) - return unmatricize(style, N, axes_codomain, (conj(axes(N, ndims(N))),)) +function unmatricize_factors(::typeof(left_null), N, axes_codomain, axes_domain) + return unmatricize(N, axes_codomain, (conj(axes(N, ndims(N))),)) end """ @@ -593,21 +533,15 @@ The output satisfies `A * Nᴴ' ≈ 0` and `Nᴴ * Nᴴ' ≈ I`. """ right_null -function right_null!!(style::MatricizeStyle, A, ndims_codomain::Val; kwargs...) - A_mat = matricize(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) +function right_null!!(A, ndims_codomain::Val; kwargs...) + A_mat = matricize(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) -end -function right_null!!(A, ndims_codomain::Val; kwargs...) - return right_null!!(MatricizeStyle(A), A, ndims_codomain; kwargs...) + return unmatricize(Nᴴ, (axes(Nᴴ, 1),), axes_domain) end -function unmatricize_factors( - ::typeof(right_null), style::MatricizeStyle, Nᴴ, - axes_codomain, axes_domain - ) - return unmatricize(style, Nᴴ, (axes(Nᴴ, 1),), axes_domain) +function unmatricize_factors(::typeof(right_null), Nᴴ, axes_codomain, axes_domain) + return unmatricize(Nᴴ, (axes(Nᴴ, 1),), axes_domain) end """ @@ -658,11 +592,8 @@ invsqrth_safe for f in (:sqrth_safe, :invsqrth_safe) @eval begin - function unmatricize_factors( - ::typeof($f), style::MatricizeStyle, P_mat, - axes_codomain, axes_domain - ) - return unmatricize(style, P_mat, axes_codomain, axes_domain) + function unmatricize_factors(::typeof($f), P_mat, axes_codomain, axes_domain) + return unmatricize(P_mat, axes_codomain, axes_domain) end end end @@ -680,11 +611,8 @@ See also `MatrixAlgebraKit.project_hermitian`. """ project_hermitian -function unmatricize_factors( - ::typeof(project_hermitian), style::MatricizeStyle, H_mat, - axes_codomain, axes_domain - ) - return unmatricize(style, H_mat, axes_codomain, axes_domain) +function unmatricize_factors(::typeof(project_hermitian), H_mat, axes_codomain, axes_domain) + return unmatricize(H_mat, axes_codomain, axes_domain) end """ @@ -707,13 +635,10 @@ See also [`MatrixAlgebra.sqrth_invsqrth_safe`](@ref). """ sqrth_invsqrth_safe -function unmatricize_factors( - ::typeof(sqrth_invsqrth_safe), style::MatricizeStyle, F, - axes_codomain, axes_domain - ) +function unmatricize_factors(::typeof(sqrth_invsqrth_safe), F, axes_codomain, axes_domain) P_mat, Pinv_mat = F - return unmatricize(style, P_mat, axes_codomain, axes_domain), - unmatricize(style, Pinv_mat, axes_codomain, axes_domain) + return unmatricize(P_mat, axes_codomain, axes_domain), + unmatricize(Pinv_mat, axes_codomain, axes_domain) end """ @@ -750,47 +675,29 @@ true function one 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 +# matricization directly when the array type 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) +function one!(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) - ) + if is_output_view(matricizeop, identity, A, perm_codomain, perm_domain) + MatrixAlgebra.one!(matricizeopview(identity, A, perm_codomain, perm_domain)) return A end - A_mat = matricizeopcopy(style, identity, A, perm_codomain, perm_domain) + A_mat = matricizeopcopy(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) - return one!(MatricizeStyle(A), A, ndims_codomain) + return unmatricize!(A, A_mat, ndims_codomain) end # 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, perm_codomain, perm_domain) - return one(MatricizeStyle(A), A, perm_codomain, perm_domain) -end -function one(style::MatricizeStyle, A, ndims_codomain::Val) - return one(style, A, identitybiperm(ndims_codomain, Val(ndims(A)))...) + A_perm = bipermutedims(A, perm_codomain, perm_domain) + return one!(A_perm, Val(length(perm_codomain))) end 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 - ) - perm_codomain, perm_domain = - biperm(Tuple.((labels_A, labels_codomain, labels_domain))...) - return one(style, A, perm_codomain, perm_domain) + return one(A, identitybiperm(ndims_codomain, Val(ndims(A)))...) end function one(A, labels_A, labels_codomain, labels_domain) perm_codomain, perm_domain = diff --git a/src/matricize.jl b/src/matricize.jl index 757dad8f..164b4d4d 100644 --- a/src/matricize.jl +++ b/src/matricize.jl @@ -1,13 +1,5 @@ using LinearAlgebra: Diagonal -# ===================================== MatricizeStyle ====================================== -abstract type MatricizeStyle end - -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 ======================================== """ @@ -70,26 +62,26 @@ function bipermutedims!( end # ===================================== matricize ======================================== -# A style implements four hooks, all taking the operation, the array and the bipermutation: +# An array type 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 +# `allocate_output(matricizeop, op, a, pc, pd)` the matrix destination +# `matricizeop!(dest, op, a, pc, pd)` write the matricization into it +# `matricizeopview(op, a, pc, pd)` partial: the aliasing form +# `is_output_view(matricizeop, 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`. +# storage the caller owns. `matricizeop` returns the view where the array type 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 +# needs the fused axes, which only the array type 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 +# `matricizeopcopy` is itself an overload point for an array type 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 type 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`. @@ -101,17 +93,14 @@ matrix representing `op.(permutedims(a, (perm_codomain..., perm_domain...)))` wi 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. +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 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) + is_output_view(matricizeop, op, a, perm_codomain, perm_domain) && + return matricizeopview(op, a, perm_codomain, perm_domain) + return matricizeopcopy(op, a, perm_codomain, perm_domain) end """ @@ -122,89 +111,61 @@ end function matricize(a, perm_codomain, perm_domain) return matricizeop(identity, a, perm_codomain, perm_domain) end -function matricize(style::MatricizeStyle, a, perm_codomain, perm_domain) - return matricizeop(style, identity, a, perm_codomain, perm_domain) -end # Split-only convenience: matricize after `ndims_codomain` dimensions without permuting. Sugar over -# the bipermutation forms, not a dispatch tier. A style implements the hooks above and never these, -# which is what keeps the copy path from recursing back through the router. +# the bipermutation forms, not a dispatch tier. An array type 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 matricize(style::MatricizeStyle, a, ndims_codomain::Val) - return matricize(style, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) -end function matricizeop(op, a, ndims_codomain::Val) return matricizeop(op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) end -function matricizeop(style::MatricizeStyle, op, a, ndims_codomain::Val) - return matricizeop(style, op, a, identitybiperm(ndims_codomain, Val(ndims(a)))...) -end # Total: always fresh storage the caller owns. function matricizeopcopy(op, a, perm_codomain, perm_domain) - return matricizeopcopy(MatricizeStyle(a), op, a, perm_codomain, perm_domain) -end -function matricizeopcopy(style::MatricizeStyle, op, a, perm_codomain, perm_domain) check_biperm(a, perm_codomain, perm_domain) - dest = allocate_output(matricizeop, style, op, a, perm_codomain, perm_domain) - return matricizeop!(dest, style, op, a, perm_codomain, perm_domain) + dest = allocate_output(matricizeop, op, a, perm_codomain, perm_domain) + return matricizeop!(dest, op, a, perm_codomain, perm_domain) end # 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)) - ) +function matricizeopview(op, a, perm_codomain, perm_domain) + return throw(MethodError(matricizeopview, (op, a, perm_codomain, perm_domain))) end -# 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)) - ) +# Required of every array type: write the matricization of `a` into `dest`. +function matricizeop!(dest, op, a, perm_codomain, perm_domain) + return throw(MethodError(matricizeop!, (dest, op, a, perm_codomain, perm_domain))) end -# Required of every style: the matrix destination `matricizeop!` writes into. -function allocate_output( - ::typeof(matricizeop), style::MatricizeStyle, op, a, perm_codomain, perm_domain - ) +# Required of every array type: the matrix destination `matricizeop!` writes into. +function allocate_output(::typeof(matricizeop), op, a, perm_codomain, perm_domain) return throw( - MethodError( - allocate_output, - (matricizeop, style, op, a, perm_codomain, perm_domain) - ) + MethodError(allocate_output, (matricizeop, op, a, perm_codomain, perm_domain)) ) end # ================================== is_output_view ====================================== -# `true` iff `matricizeop(style, op, a, perm_codomain, perm_domain)` shares `a`'s memory, so that +# `true` iff `matricizeop(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 - ) +# `output_axes`), so the predicate's arguments are exactly the call's arguments. The untyped +# default declares nothing, which is always safe. +function is_output_view(::typeof(matricizeop), op, a, perm_codomain, perm_domain) return false end # 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 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) - return throw(MethodError(unmatricize, (style, m, axes_codomain, axes_domain))) -end -function unmatricize(m, axes_codomain, axes_domain) - return unmatricize(MatricizeStyle(m), m, axes_codomain, axes_domain) -end +# matrix type that stores its domain axes dualized re-dualizes them with `conj` in its overload +# (a no-op on a dense axis). This is the primary overload point for a new matrix type, dispatched +# on `m`. 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 end # Split `axes` into its codomain and domain groups like `bipartition`, but present the domain # group codomain-facing (un-dualized) with `conj`, the convention `unmatricize` and `similar_map` @@ -215,59 +176,66 @@ function bipartition_axes(t::Tuple, split...) return axes_codomain, conj.(axes_domain) end +# `a_dest = β * a_dest + α * unmatricize(m)`, scattered across the bipermutation in one pass. # 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!( +# the same forward bipermutation to both. Dispatched on `a_dest`: a wrapper type scatters into its +# parent by overloading this form, and `unmatricize!` is the `(1, 0)` case. +function unmatricizeadd!( a_dest, m, - perm_codomain, perm_domain - ) - return unmatricize!(MatricizeStyle(m), a_dest, m, perm_codomain, perm_domain) -end -function unmatricize!( - style::MatricizeStyle, a_dest, m, - perm_codomain, perm_domain + perm_codomain, perm_domain, + α::Number, β::Number ) 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), biperm_src) - a_perm = unmatricize(style, m, axes_codomain, axes_domain) + a_perm = unmatricize(m, axes_codomain, axes_domain) biperm_dest = BiTuple(Tuple(invperm(biperm_src)), Val(length_codomain(biperm_src))) - return bipermutedims!(a_dest, a_perm, biperm_dest) + return bipermutedimsopadd!( + a_dest, + identity, + a_perm, + biperm_dest.t1, + biperm_dest.t2, + α, + β + ) +end + +function unmatricize!(a_dest, m, perm_codomain, perm_domain) + return unmatricizeadd!(a_dest, m, perm_codomain, perm_domain, true, false) end # 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 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) - return unmatricize!( - style, a_dest, m, identitybiperm(ndims_codomain, Val(ndims(a_dest)))... - ) -end function unmatricize!(a_dest, m, ndims_codomain::Val) - return unmatricize!(MatricizeStyle(a_dest), a_dest, m, ndims_codomain) + return unmatricize!(a_dest, m, identitybiperm(ndims_codomain, Val(ndims(a_dest)))...) end -# Defaults to ReshapeMatricize, a simple reshape -struct ReshapeMatricize <: MatricizeStyle end -MatricizeStyle(::Type{<:AbstractArray}) = ReshapeMatricize() -# 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`. +# ================================ dense (reshape) hooks =================================== +# The view must be a matrix that BLAS and LAPACK handle efficiently, not merely one that shares +# memory. A `DenseArray` (`Array`, GPU arrays) reshapes to one of its own kind, so it declares the +# view; any other `AbstractArray` reshapes to a `ReshapedArray` wrapper that `mul!` sends to +# generic matmul, so it takes the owned copy instead and lands on BLAS. The 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`. The copy hooks below stay generic, since +# `similar` plus a permuted add is a correct owned matricization for any permutable array. function is_output_view( - ::typeof(matricizeop), ::ReshapeMatricize, op, a, perm_codomain, perm_domain + ::typeof(matricizeop), op, a::DenseArray, perm_codomain, perm_domain ) return op === identity && isidentitybiperm(perm_codomain, perm_domain) end -function matricizeopview(::ReshapeMatricize, op, a, perm_codomain, perm_domain) +function matricizeopview(op, a::DenseArray, 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 allocate_output( - ::typeof(matricizeop), ::ReshapeMatricize, op, a, perm_codomain, perm_domain + ::typeof(matricizeop), op, a::AbstractArray, perm_codomain, perm_domain ) T = Base.promote_op(op, eltype(a)) size_codomain = map(i -> size(a, i), perm_codomain) @@ -276,7 +244,7 @@ function allocate_output( 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) +function matricizeop!(dest, op, a::AbstractArray, 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) @@ -295,7 +263,7 @@ function check_input(::typeof(unmatricize), m, axes_codomain, axes_domain) end # A dense reshape ignores the codomain/domain split: it just reshapes to the concatenated axes. # `conj` re-dualizes the codomain-facing `axes_domain` into stored form, a no-op on a dense axis. -function unmatricize(style::ReshapeMatricize, m, axes_codomain, axes_domain) +function unmatricize(m::AbstractMatrix, axes_codomain, axes_domain) check_input(unmatricize, m, axes_codomain, axes_domain) return reshape(m, (axes_codomain..., conj.(axes_domain)...)) end diff --git a/src/matrixfunctions.jl b/src/matrixfunctions.jl index a128f22f..e6052321 100644 --- a/src/matrixfunctions.jl +++ b/src/matrixfunctions.jl @@ -36,48 +36,19 @@ const MATRIX_FUNCTIONS = [ # the eager `bipermutedims` copy at the identity bipermutation. for f in MATRIX_FUNCTIONS @eval begin - function $f(style::MatricizeStyle, a, ndims_codomain::Val; kwargs...) - return $f( - style, a, identitybiperm(ndims_codomain, Val(ndims(a)))...; kwargs... - ) - end function $f(a, ndims_codomain::Val; kwargs...) - return $f(MatricizeStyle(a), a, ndims_codomain; kwargs...) + return $f(a, identitybiperm(ndims_codomain, Val(ndims(a)))...; kwargs...) end - - function $f( - style::MatricizeStyle, a, - perm_codomain, perm_domain; - kwargs... - ) - a_mat = matricize(style, a, perm_codomain, perm_domain) + function $f(a, perm_codomain, perm_domain; kwargs...) + a_mat = matricize(a, perm_codomain, perm_domain) axes_codomain, axes_domain = bipartition_axes( map(i -> axes(a, i), (perm_codomain..., perm_domain...)), Val(length(perm_codomain)) ) fa_mat = Base.$f(a_mat; kwargs...) - return unmatricize(style, fa_mat, axes_codomain, axes_domain) - end - function $f( - a, - perm_codomain, perm_domain; - kwargs... - ) - return $f(MatricizeStyle(a), a, perm_codomain, perm_domain; kwargs...) + return unmatricize(fa_mat, axes_codomain, axes_domain) end - - function $f( - style::MatricizeStyle, a, - labels_a, labels_codomain, labels_domain; kwargs... - ) - perm_codomain, perm_domain = - biperm(Tuple.((labels_a, labels_codomain, labels_domain))...) - return $f(style, a, perm_codomain, perm_domain; kwargs...) - end - function $f( - a, - labels_a, labels_codomain, labels_domain; kwargs... - ) + function $f(a, labels_a, labels_codomain, labels_domain; kwargs...) perm_codomain, perm_domain = biperm(Tuple.((labels_a, labels_codomain, labels_domain))...) return $f(a, perm_codomain, perm_domain; kwargs...) diff --git a/test/Project.toml b/test/Project.toml index a2924b93..c07425b2 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -35,7 +35,7 @@ Random = "1.10" SafeTestsets = "0.1" StableRNGs = "1.0.2" Suppressor = "0.2" -TensorAlgebra = "0.21" +TensorAlgebra = "0.22" TensorKit = "0.17" TensorOperations = "5.1.4" Test = "1.10" diff --git a/test/test_basics.jl b/test/test_basics.jl index 1903aab2..d640718f 100644 --- a/test/test_basics.jl +++ b/test/test_basics.jl @@ -2,7 +2,7 @@ import TensorAlgebra using StableRNGs: StableRNG using TensorAlgebra: BiTuple, ContractAlgorithm, bipermutedims, bipermutedims!, contract, contract!, contractadd!, contractalign, length_codomain, length_domain, matricize, - unmatricize, unmatricize! + unmatricize, unmatricize!, unmatricizeadd! using TensorOperations: TensorOperations using Test: @test, @test_broken, @test_throws, @testset @@ -146,6 +146,16 @@ TensorAlgebra.label_type(::Type{OptInLabel}) = Int unmatricize!(a, m, invperm_codomain, invperm_domain) @test a ≈ a1 + # `unmatricizeadd!` accumulates during the scatter; `unmatricize!` is its `(1, 0)` case. + α, β = elt(2), elt(-3) + a = randn(elt, size(a1)) + a_expected = α * a1 + β * a + unmatricizeadd!(a, m, invperm_codomain, invperm_domain, α, β) + @test a ≈ a_expected + a = fill(elt(NaN), size(a1)) + unmatricizeadd!(a, m, invperm_codomain, invperm_domain, α, false) + @test a ≈ α * a1 + a = unmatricize(reshape(a0, 1, 120), (), axes0) @test eltype(a) === elt @test a ≈ a0 diff --git a/test/test_diagonal.jl b/test/test_diagonal.jl index fd751107..8b49e32a 100644 --- a/test/test_diagonal.jl +++ b/test/test_diagonal.jl @@ -41,13 +41,13 @@ using Test: @test, @test_throws, @testset end @testset "matricize(1, 1) is the identity reshape" begin - m = TensorAlgebra.matricize(TensorAlgebra.ReshapeMatricize(), d, (1,), (2,)) + m = TensorAlgebra.matricize(d, (1,), (2,)) @test m === d end @testset "unmatricize round-trips a Diagonal on its own {1,1} axes" begin ax = axes(d, 1) - back = TensorAlgebra.unmatricize(TensorAlgebra.ReshapeMatricize(), d, (ax,), (ax,)) + back = TensorAlgebra.unmatricize(d, (ax,), (ax,)) @test back === d end @@ -55,18 +55,14 @@ using Test: @test, @test_throws, @testset d4 = Diagonal(elt[1, 2, 3, 4]) axes_codomain = (Base.OneTo(2), Base.OneTo(2)) axes_domain = (Base.OneTo(4),) - t = TensorAlgebra.unmatricize( - TensorAlgebra.ReshapeMatricize(), d4, axes_codomain, axes_domain - ) + t = TensorAlgebra.unmatricize(d4, axes_codomain, axes_domain) @test !(t isa Diagonal) @test t == reshape(Array(d4), 2, 2, 4) end @testset "unmatricize errors on a mismatched {1,1} split" begin wrong = Base.OneTo(length(diag(d)) + 1) - @test_throws DimensionMismatch TensorAlgebra.unmatricize( - TensorAlgebra.ReshapeMatricize(), d, (wrong,), (wrong,) - ) + @test_throws DimensionMismatch TensorAlgebra.unmatricize(d, (wrong,), (wrong,)) end @testset "matrix functions preserve Diagonal" begin @@ -80,51 +76,42 @@ using Test: @test, @test_throws, @testset 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) + ref = TensorAlgebra.matricizeop(op, dz, pc, pd) + if TensorAlgebra.is_output_view(TensorAlgebra.matricizeop, op, dz, pc, pd) + m = TensorAlgebra.matricizeopview(op, dz, pc, pd) @test Base.mightalias(m, dz) @test m == ref end - m_copy = TensorAlgebra.matricizeopcopy(style, op, dz, pc, pd) + m_copy = TensorAlgebra.matricizeopcopy(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,) + TensorAlgebra.matricizeop, + identity, + dz, + (1,), + (2,) ) + @test !TensorAlgebra.is_output_view(TensorAlgebra.matricizeop, 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 TensorAlgebra.unmatricize!(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 @@ -135,9 +122,6 @@ using Test: @test, @test_throws, @testset 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 diff --git a/test/test_exports.jl b/test/test_exports.jl index 36e952e1..f9f2ae26 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -32,8 +32,8 @@ using Test: @test, @testset :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, + :lq_full, :matricize, :MatricizeContract, :matricizeop, + :matricizeop!, :matricizeopcopy, :matricizeopview, :MATRIX_FUNCTIONS, :ndims, :ndims_codomain, :ndims_domain, :one, :ones_map, :operation, :output_axes, :PermutedDims, :permuteddims, :permutedims, :permutedims!, @@ -45,7 +45,8 @@ using Test: @test, @testset :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!, + :unmatricize!, :unmatricize_factors, :unmatricizeadd!, :unproject, + :unscaled, :zero!, :zeros_map, ] ) diff --git a/test/test_factorizations.jl b/test/test_factorizations.jl index b6bfe7db..6e022f92 100644 --- a/test/test_factorizations.jl +++ b/test/test_factorizations.jl @@ -343,12 +343,9 @@ end @test TensorAlgebra.matricize(Id, splitperms(Id, 2)...) ≈ I - # `Val`, perm, and label entries agree, as do the style-explicit spellings of each. - style = TensorAlgebra.MatricizeStyle(A) + # `Val`, perm, and label entries agree. @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 @@ -367,10 +364,6 @@ end @test Cret === C @test TensorAlgebra.matricize(C, splitperms(C, 2)...) ≈ I @test C ≈ TensorAlgebra.one(A, Val(2)) - Cstyle = randn(T, 2, 3, 2, 3) - @test TensorAlgebra.one!(TensorAlgebra.MatricizeStyle(Cstyle), Cstyle, Val(2)) === - Cstyle - @test TensorAlgebra.matricize(Cstyle, splitperms(Cstyle, 2)...) ≈ I # `unmatricize!` scatters a fused matrix back into an existing array. D = randn(T, 2, 3, 2, 3) @@ -456,44 +449,23 @@ module FactorizationMatricizeTestUtils function Base.getindex(a::AliasingArray{<:Any, N}, I::Vararg{Int, N}) where {N} return a.parent[I...] end - struct AliasingMatricize <: TA.MatricizeStyle end - TA.MatricizeStyle(::Type{<:AliasingArray}) = AliasingMatricize() - # Delegate every hook to the dense style on the unwrapped parent, so the matricization + # Delegate every hook to the dense hooks 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 + ::typeof(TA.matricizeop), op, a::AliasingArray, perm_codomain, perm_domain ) + return TA.is_output_view(TA.matricizeop, op, a.parent, perm_codomain, perm_domain) end - function TA.matricizeopview( - ::AliasingMatricize, op, a, perm_codomain, perm_domain - ) - return TA.matricizeopview( - TA.ReshapeMatricize(), op, unwrap(a), perm_codomain, perm_domain - ) + function TA.matricizeopview(op, a::AliasingArray, perm_codomain, perm_domain) + return TA.matricizeopview(op, a.parent, 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 + ::typeof(TA.matricizeop), op, a::AliasingArray, perm_codomain, perm_domain ) + return TA.allocate_output(TA.matricizeop, op, a.parent, 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( - TA.unmatricize(TA.ReshapeMatricize(), m, axes_codomain, axes_domain) - ) + function TA.matricizeop!(dest, op, a::AliasingArray, perm_codomain, perm_domain) + return TA.matricizeop!(dest, op, a.parent, perm_codomain, perm_domain) end end using .FactorizationMatricizeTestUtils: AliasingArray diff --git a/test/test_matricize.jl b/test/test_matricize.jl index 1b2dd58c..070bfb88 100644 --- a/test/test_matricize.jl +++ b/test/test_matricize.jl @@ -1,10 +1,10 @@ using StableRNGs: StableRNG -using TensorAlgebra: TensorAlgebra, ReshapeMatricize, is_output_view, matricize, - matricizeop, matricizeop!, matricizeopcopy, matricizeopview +using TensorAlgebra: TensorAlgebra, 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. -struct DummyMatricize <: TensorAlgebra.MatricizeStyle end +# Not an `AbstractArray`, so none of the dense hooks apply: checks the always-safe fallbacks. +struct DummyArray end # Ground-truth matricization: permute into `(codomain..., domain...)` order, then reshape. function matricize_ref(a, perm_codomain, perm_domain) @@ -43,25 +43,30 @@ end @testset "is_output_view" begin a = randn(StableRNG(321), 2, 3, 4) - style = ReshapeMatricize() # 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), ()) + @test is_output_view(matricizeop, identity, a, (1,), (2, 3)) + @test is_output_view(matricizeop, identity, a, (), (1, 2, 3)) + @test is_output_view(matricizeop, identity, a, (1, 2, 3), ()) # 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,)) + @test !is_output_view(matricizeop, identity, a, (2, 3), (1,)) + @test !is_output_view(matricizeop, 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)) + @test !is_output_view(matricizeop, conj, a, (1,), (2, 3)) - # A generic style declares nothing (fail-safe default). - @test !is_output_view(matricizeop, DummyMatricize(), identity, a, (1,), (2, 3)) + # A type without hooks declares nothing (fail-safe default) and has no aliasing form. + @test !is_output_view(matricizeop, identity, DummyArray(), (1,), (2,)) + @test_throws MethodError matricizeopview(identity, DummyArray(), (1,), (2,)) + # Only a `DenseArray` declares the reshape a view; a strided view of one does not, even in + # codomain-then-domain order, and still matricizes correctly through the copy path. + a_view = view(a, :, :, :) + @test !is_output_view(matricizeop, identity, a_view, (1,), (2, 3)) + @test matricize(a_view, (1,), (2, 3)) == matricize(a, (1,), (2, 3)) # Writes to the shared matricization are writes to `a`. - m = matricizeopview(style, identity, a, (1,), (2, 3)) + m = matricizeopview(identity, a, (1,), (2, 3)) @test m == matricize_ref(a, (1,), (2, 3)) m[1, 1] = 42 @test a[1, 1, 1] == 42 @@ -70,37 +75,36 @@ end @testset "is_output_view coherence" begin rng = StableRNG(11) a = randn(rng, 2, 3, 4) - style = ReshapeMatricize() # A declared share means `matricizeopview` (and so `matricize`) aliases `a`, while # `matricizeopcopy` never does. for K in 0:3 pc = ntuple(identity, K) pd = ntuple(i -> K + i, 3 - K) - if is_output_view(matricizeop, style, identity, a, pc, pd) - m = matricizeopview(style, identity, a, pc, pd) + if is_output_view(matricizeop, identity, a, pc, pd) + m = matricizeopview(identity, a, pc, pd) @test Base.mightalias(m, a) - @test matricize(style, a, pc, pd) == m + @test matricize(a, pc, pd) == m end - m_copy = matricizeopcopy(style, identity, a, pc, pd) + m_copy = matricizeopcopy(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 = matricizeopcopy(style, identity, a, pc, pd) + m = matricizeopcopy(identity, a, pc, pd) @test m ≈ matricize_ref(a, pc, pd) @test !Base.mightalias(m, a) end - @test_throws ArgumentError matricizeopcopy(style, identity, a, (1,), (2,)) + @test_throws ArgumentError matricizeopcopy(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) + dest = TensorAlgebra.allocate_output(matricizeop, op, a, pc, pd) @test size(dest) == size(matricize_ref(a, pc, pd)) - matricizeop!(dest, style, op, a, pc, pd) + matricizeop!(dest, op, a, pc, pd) @test dest ≈ op.(matricize_ref(a, pc, pd)) - @test dest ≈ matricizeopcopy(style, op, a, pc, pd) + @test dest ≈ matricizeopcopy(op, a, pc, pd) end end diff --git a/test/test_matricizehooks.jl b/test/test_matricizehooks.jl new file mode 100644 index 00000000..fdce14d9 --- /dev/null +++ b/test/test_matricizehooks.jl @@ -0,0 +1,52 @@ +using LinearAlgebra: I +using TensorAlgebra: TensorAlgebra as TA, MatricizeContract +using Test: @test, @testset + +module MatricizeHooksTestUtils + using TensorAlgebra: TensorAlgebra as TA + struct MyArray{T, N, A <: AbstractArray{T, N}} <: AbstractArray{T, N} + parent::A + end + # Minimal hooks so a round trip (`one!`) can run through the wrapper. All of them dispatch on + # `MyArray`, so a path that dropped the wrapper and re-derived the hooks from the plain fused + # matrix would miss them and error. + function TA.is_output_view( + ::typeof(TA.matricizeop), op, a::MyArray, perm_codomain, perm_domain + ) + return false + end + function TA.allocate_output( + ::typeof(TA.matricizeop), op, a::MyArray, perm_codomain, perm_domain + ) + return TA.allocate_output(TA.matricizeop, op, a.parent, perm_codomain, perm_domain) + end + function TA.matricizeop!(dest, op, a::MyArray, perm_codomain, perm_domain) + return TA.matricizeop!(dest, op, a.parent, perm_codomain, perm_domain) + end + function TA.unmatricizeadd!( + a_dest::MyArray, m, perm_codomain, perm_domain, α::Number, β::Number + ) + TA.unmatricizeadd!(a_dest.parent, m, perm_codomain, perm_domain, α, β) + return a_dest + end +end +using .MatricizeHooksTestUtils: MyArray + +@testset "MatricizeContract is the dense default" begin + a1 = randn(2, 2) + a2 = MyArray(randn(2, 2)) + @test MatricizeContract() ≡ MatricizeContract() + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a1), typeof(a1), typeof(a1)}) ≡ + MatricizeContract() + @test TA.default_algorithm(TA.contract!, Tuple{typeof(a2), typeof(a2), typeof(a2)}) ≡ + MatricizeContract() +end + +@testset "the hooks thread through the unfold" begin + # `one!` folds through the wrapper's hooks and must unfold through them too, not through + # hooks re-derived from the fused matrix (here a plain `Matrix`, whose dense `unmatricize!` + # would not know how to scatter into a `MyArray`). + A = MyArray(randn(3, 3)) + TA.one!(A, Val(1)) + @test A.parent ≈ I +end diff --git a/test/test_matricizestyle.jl b/test/test_matricizestyle.jl deleted file mode 100644 index 5dd25b17..00000000 --- a/test/test_matricizestyle.jl +++ /dev/null @@ -1,75 +0,0 @@ -using LinearAlgebra: I -using TensorAlgebra: - TensorAlgebra as TA, MatricizeContract, 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() - # Minimal hooks so a round trip (`one!`) can run through the custom style. All of them - # dispatch on `MyArrayMatricize`, so a path whose style was re-derived from the plain fused - # matrix instead of threaded through would miss them and error. - function TA.is_output_view( - ::typeof(TA.matricizeop), ::MyArrayMatricize, op, a, perm_codomain, perm_domain - ) - return false - end - function TA.allocate_output( - ::typeof(TA.matricizeop), ::MyArrayMatricize, op, a::MyArray, - perm_codomain, perm_domain - ) - return TA.allocate_output( - TA.matricizeop, TA.ReshapeMatricize(), op, a.parent, perm_codomain, perm_domain - ) - end - function TA.matricizeop!( - dest, ::MyArrayMatricize, op, a::MyArray, perm_codomain, perm_domain - ) - return TA.matricizeop!( - dest, TA.ReshapeMatricize(), op, a.parent, perm_codomain, perm_domain - ) - end - function TA.unmatricize!( - ::MyArrayMatricize, a_dest::MyArray, m, - perm_codomain, perm_domain - ) - TA.unmatricize!( - TA.ReshapeMatricize(), a_dest.parent, m, perm_codomain, perm_domain - ) - return a_dest - end -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_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 - # `one!` folds with the caller-supplied style and must unfold with the same style, not one - # re-derived from the fused matrix (here a plain `Matrix`, whose derived style would be - # `ReshapeMatricize` and would not know how to scatter into a `MyArray`). - A = MyArray(randn(3, 3)) - TA.one!(MyArrayMatricize(), A, Val(1)) - @test A.parent ≈ I -end diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index 6125a37c..10e3b060 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -337,7 +337,6 @@ using Test: @test, @test_throws, @testset # 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 @@ -346,8 +345,6 @@ using Test: @test, @test_throws, @testset 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 @@ -357,9 +354,6 @@ using Test: @test, @test_throws, @testset 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 @@ -367,17 +361,16 @@ using Test: @test, @test_throws, @testset 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) + if TensorAlgebra.is_output_view(TensorAlgebra.matricizeop, 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 + @test TensorAlgebra.matricizeopview(op, t, pc, pd) === t end - m_copy = TensorAlgebra.matricizeopcopy(style, op, t, pc, pd) + m_copy = TensorAlgebra.matricizeopcopy(op, t, pc, pd) @test m_copy !== t @test TensorAlgebra.data(m_copy) !== TensorAlgebra.data(t) @test space(m_copy) == space(ref) @@ -385,16 +378,14 @@ using Test: @test, @test_throws, @testset 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,) + TensorAlgebra.matricizeop, identity, t, (1, 2), (3,) ) @test !TensorAlgebra.is_output_view( - TensorAlgebra.matricizeop, style, identity, t, (1, 3), (2,) + TensorAlgebra.matricizeop, identity, t, (1, 3), (2,) ) dest = similar(t) - @test TensorAlgebra.unmatricize!( - style, dest, matricize(style, t, (1, 2), (3,)), Val(2) - ) === dest + @test TensorAlgebra.unmatricize!(dest, matricize(t, (1, 2), (3,)), Val(2)) === dest @test dest ≈ t end