Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "TensorAlgebra"
uuid = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a"
version = "0.21.0"
version = "0.21.1"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
1 change: 1 addition & 0 deletions ext/TensorAlgebraTensorKitExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion src/TensorAlgebra.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
19 changes: 17 additions & 2 deletions src/projectto.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion test/test_exports.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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!,
Expand Down
14 changes: 11 additions & 3 deletions test/test_projectto.jl
Original file line number Diff line number Diff line change
@@ -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)
Expand Down Expand Up @@ -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))
Expand Down
16 changes: 12 additions & 4 deletions test/test_tensorkitext.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Loading