diff --git a/Project.toml b/Project.toml index 3fe422fc..29a6dcc6 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "TensorAlgebra" uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a" -version = "0.21.0" +version = "0.21.1" authors = ["ITensor developers and contributors"] [workspace] diff --git a/ext/TensorAlgebraTensorKitExt.jl b/ext/TensorAlgebraTensorKitExt.jl index a17957e1..608f8d4c 100644 --- a/ext/TensorAlgebraTensorKitExt.jl +++ b/ext/TensorAlgebraTensorKitExt.jl @@ -23,6 +23,7 @@ TensorAlgebra.size(t::AbstractTensorMap, i::Int) = dim(space(t, i)) TensorAlgebra.size(t::AbstractTensorMap) = ntuple(i -> dim(space(t, i)), numind(t)) # A `TensorMap` stores its codomain/domain split, so its codomain rank is `numout`. TensorAlgebra.ndims_codomain(t::AbstractTensorMap) = numout(t) +TensorAlgebra.has_bipartition(::AbstractTensorMap) = true # `t[]` on a rank-0 `TensorMap` requires a trivial sector type; `TensorKit.scalar` is the # general spelling. diff --git a/src/TensorAlgebra.jl b/src/TensorAlgebra.jl index 63968afb..92887326 100644 --- a/src/TensorAlgebra.jl +++ b/src/TensorAlgebra.jl @@ -5,7 +5,7 @@ export contract, contract!, contractalign, dual, isdual, MatrixAlgebra if VERSION >= v"1.11.0-DEV.469" eval( Meta.parse( - "public AbstractAlgorithm, add!, AddBroadcasted, addends, allocate_output, allocate_project, arguments, axes, bipartition, bipartition_axes, biperm, bipermutedims, bipermutedims!, bipermutedimsopadd!, cat_axis, cat_similar, check_input, concatenate, concatenate!, ConjBroadcasted, contractadd!, ContractAlgorithm, contractopadd!, contractperm, contractperm!, contractpermadd!, contractpermalign, contractpermopadd!, data, datatype, default_algorithm, dims2cat, directsum, eig_full, eig_trunc, eig_vals, eigh_full, eigh_trunc, eigh_vals, fill_map, flattenlinear, infer_aux_space, invsqrth_safe, is_output_view, is_projected, isidentitybiperm, label_type, left_null, left_orth, left_polar, LinearBroadcasted, linearbroadcasted, lq_compact, lq_full, matricize, MatricizeContract, matricizeop, matricizeop!, matricizeopcopy, matricizeopview, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, 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, MatricizeStyle, MATRIX_FUNCTIONS, ndims, ndims_codomain, ndims_domain, one, ones_map, operation, output_axes, PermutedDims, permuteddims, permutedims, permutedims!, permutedimsadd!, permutedimsop, permutedimsopadd!, project, project!, project_aux, project_hermitian, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, right_orth, right_polar, scalar, scale!, ScaledBroadcasted, select_algorithm, similar_map, size, sqrth_invsqrth_safe, sqrth_safe, sum, svd_compact, svd_full, svd_trunc, svd_vals, TensorOperationsContract, to_range, tr, trivialrange, tryflattenlinear, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, ungrade, unmatricize, unmatricize!, unmatricize_factors, unproject, unscaled, zero!, zeros_map" ) ) end diff --git a/src/projectto.jl b/src/projectto.jl index 59c53862..048a5159 100644 --- a/src/projectto.jl +++ b/src/projectto.jl @@ -80,13 +80,28 @@ The codomain and domain ranks of `a`'s intrinsic split, for a type that stores o no intrinsic split, so it defaults to all codomain and an empty domain, and a type that does store one overloads `ndims_codomain` (a `TensorMap` returns `numout`). -Only `ndims_codomain` needs overloading: `ndims_domain` is whatever rank is left over, so the two -agree by construction. +`ndims_domain` never needs overloading: it is whatever rank is left over, so the two agree by +construction. A type that does overload `ndims_codomain` should also overload +[`has_bipartition`](@ref). """ ndims_codomain(a) = ndims(a) @doc (@doc ndims_codomain) ndims_domain(a) = ndims(a) - ndims_codomain(a) +""" + TensorAlgebra.has_bipartition(a) -> Bool + +Whether `a`'s type stores an intrinsic codomain/domain split, so that +[`ndims_codomain`](@ref) reports a split `a` genuinely carries rather than the all-codomain +fallback. A `TensorMap` returns `true`, an array `false`; a type that overloads +`ndims_codomain` should also overload this. + +Lets a consumer tell a genuinely all-codomain `a` apart from one with no notion of a split, +which is what validating a caller-supplied split against `a` needs: there is nothing to +validate it against unless `a` stores one. +""" +has_bipartition(a) = false + """ is_projected(dest, src, ndims_codomain::Val; kwargs...) -> Bool is_projected(dest, src; kwargs...) -> Bool diff --git a/test/test_exports.jl b/test/test_exports.jl index 90a3c2f9..36e952e1 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -28,7 +28,8 @@ using Test: @test, @testset :datatype, :default_algorithm, :dims2cat, :directsum, :eig_full, :eig_trunc, :eig_vals, :eigh_full, :eigh_trunc, :eigh_vals, :fill_map, - :flattenlinear, :infer_aux_space, :invsqrth_safe, :is_output_view, + :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!, diff --git a/test/test_projectto.jl b/test/test_projectto.jl index c83acb43..942c5168 100644 --- a/test/test_projectto.jl +++ b/test/test_projectto.jl @@ -1,6 +1,6 @@ -using TensorAlgebra: TensorAlgebra, is_projected, project, project!, project_aux, - projectto!, tryproject, tryproject_aux, unchecked_project, unchecked_project_aux, - unproject +using TensorAlgebra: TensorAlgebra, has_bipartition, is_projected, ndims_codomain, + ndims_domain, project, project!, project_aux, projectto!, tryproject, tryproject_aux, + unchecked_project, unchecked_project_aux, unproject using Test: @test, @test_throws, @testset const elts = (Float32, Float64, ComplexF32, ComplexF64) @@ -83,6 +83,14 @@ end @test tryproject(raw, (Base.OneTo(2),), (Base.OneTo(3),)) == raw end +@testset "the codomain/domain accessors on an array" begin + a = randn(2, 3, 4) + @test ndims_codomain(a) == 3 + @test ndims_domain(a) == 0 + # An array is all-codomain by fallback, not because it stores that split. + @test !has_bipartition(a) +end + @testset "unproject (dense default) ($T)" for T in elts raw = randn(T, 2, 3, 2, 3) cod = (Base.OneTo(2), Base.OneTo(3)) diff --git a/test/test_tensorkitext.jl b/test/test_tensorkitext.jl index e4446fdc..6125a37c 100644 --- a/test/test_tensorkitext.jl +++ b/test/test_tensorkitext.jl @@ -2,10 +2,10 @@ using Base.Broadcast: broadcasted using LinearAlgebra: LinearAlgebra, norm using StableRNGs: StableRNG using TensorAlgebra: TensorAlgebra, contract, contractalign, eig_full, eig_vals, eigh_full, - eigh_vals, left_null, left_orth, left_polar, lq_compact, lq_full, matricize, project, - project_aux, projectto!, qr_compact, qr_full, rand_map, randn_map, right_null, - right_orth, right_polar, similar_map, svd_compact, svd_full, svd_vals, tryflattenlinear, - tryproject, unchecked_project, unmatricize, zeros_map + eigh_vals, has_bipartition, left_null, left_orth, left_polar, lq_compact, lq_full, + matricize, project, project_aux, projectto!, qr_compact, qr_full, rand_map, randn_map, + right_null, right_orth, right_polar, similar_map, svd_compact, svd_full, svd_vals, + tryflattenlinear, tryproject, unchecked_project, unmatricize, zeros_map using TensorKit: TensorKit, @tensor, AbstractTensorMap, DiagonalTensorMap, Irrep, Rep, SU₂, TensorMap, U₁, dim, dual, fuse, isomorphism, randn, reduceddim, space, storagetype, ←, ⊗ using Test: @test, @test_throws, @testset @@ -493,3 +493,11 @@ end s = Irrep[U₁](1) @test TensorAlgebra.dual(s) == dual(s) end + +@testset "the codomain/domain accessors on a TensorMap" begin + W = Rep[U₁](0 => 1, 1 => 1) + t = zeros_map(Float64, (W, W), (W,)) + @test has_bipartition(t) + @test TensorAlgebra.ndims_codomain(t) == 2 + @test TensorAlgebra.ndims_domain(t) == 1 +end