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.1"
version = "0.22.0"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
2 changes: 1 addition & 1 deletion docs/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,4 +11,4 @@ path = ".."
Documenter = "1.8.1"
ITensorFormatter = "0.2.27"
Literate = "2.20.1"
TensorAlgebra = "0.21"
TensorAlgebra = "0.22"
2 changes: 1 addition & 1 deletion examples/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -5,4 +5,4 @@ TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a"
path = ".."

[compat]
TensorAlgebra = "0.21"
TensorAlgebra = "0.22"
14 changes: 5 additions & 9 deletions ext/TensorAlgebraTensorKitExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -241,22 +241,20 @@ 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 &&
TensorAlgebra.isidentitybiperm(perm_codomain, perm_domain) &&
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
Expand All @@ -265,15 +263,15 @@ 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(
TensorAlgebra.permutedimsop, op, t, perm_codomain, perm_domain
)
end
function TensorAlgebra.matricizeop!(
dest::AbstractTensorMap, ::TensorKitMatricize, op,
dest::AbstractTensorMap, op,
t::AbstractTensorMap, perm_codomain, perm_domain
)
return TensorAlgebra.bipermutedimsopadd!(
Expand All @@ -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 ||
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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))
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, 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
Expand Down
18 changes: 7 additions & 11 deletions src/contract/contract.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
44 changes: 15 additions & 29 deletions src/contract/contract_matricize.jl
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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
10 changes: 1 addition & 9 deletions src/contract/contractalgorithm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
20 changes: 12 additions & 8 deletions src/diagonal.jl
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -43,17 +43,23 @@ 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
# A `{1,1}` unmatricize (one codomain axis, one domain axis) is the endomorphism identity: the
# result stays `Diagonal`, so return `m` directly. The generic `check_input(unmatricize, ...)`
# validates the axis lengths against `m`'s size.
function unmatricize(
::ReshapeMatricize, m::Diagonal,
m::Diagonal,
axes_codomain::Tuple{<:AbstractUnitRange}, axes_domain::Tuple{<:AbstractUnitRange}
)
check_input(unmatricize, m, axes_codomain, axes_domain)
Expand All @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/directsum.jl
Original file line number Diff line number Diff line change
@@ -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...)
Loading
Loading