Skip to content
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "ITensorBase"
uuid = "4795dd04-0d67-49bb-8f44-b89c448a1dc7"
version = "0.15.2"
version = "0.15.3"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
2 changes: 1 addition & 1 deletion src/ITensorBase.jl
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ export AbstractNamedTensor, NamedTensor, AbstractITensor, ITensor, Index,
if VERSION >= v"1.11.0-DEV.469"
eval(
Meta.parse(
"public @names, IndexName, name, names, setname, space, unnamed, unnamedtype, decoration, emptytags, gettag, gettags, hastag, plev, settags, tags, unsettags"
"public @names, IndexName, mulopadd!, name, names, setname, space, unnamed, unnamedtype, decoration, emptytags, gettag, gettags, hastag, plev, settags, tags, unsettags"
)
)
end
Expand Down
42 changes: 34 additions & 8 deletions src/tensoralgebra.jl
Original file line number Diff line number Diff line change
Expand Up @@ -37,38 +37,64 @@ end
function LA.mul!(
a_dest::AbstractNamedTensor,
a1::AbstractNamedTensor, a2::AbstractNamedTensor,
α::Number, β::Number
α::Number, β::Number;
kwargs...
)
return mul!_namedtensor(a_dest, a1, a2, α, β)
return mul!_namedtensor(a_dest, a1, a2, α, β; kwargs...)
end
function mul!_namedtensor(
a_dest::AbstractNamedTensor,
a1::AbstractNamedTensor, a2::AbstractNamedTensor,
α::Number, β::Number
α::Number, β::Number;
kwargs...
)
TA.contractadd!(
unnamed(a_dest), names(a_dest),
unnamed(a1), names(a1),
unnamed(a2), names(a2),
α, β
α, β; kwargs...
)
return a_dest
end

function LA.mul!(
a_dest::AbstractNamedTensor,
a1::AbstractNamedTensor, a2::AbstractNamedTensor
a1::AbstractNamedTensor, a2::AbstractNamedTensor;
kwargs...
)
return mul!_namedtensor(a_dest, a1, a2)
return mul!_namedtensor(a_dest, a1, a2; kwargs...)
end
function mul!_namedtensor(
a_dest::AbstractNamedTensor,
a1::AbstractNamedTensor, a2::AbstractNamedTensor
a1::AbstractNamedTensor, a2::AbstractNamedTensor;
kwargs...
)
TA.contract!(
unnamed(a_dest), names(a_dest),
unnamed(a1), names(a1),
unnamed(a2), names(a2)
unnamed(a2), names(a2); kwargs...
)
return a_dest
end

"""
mulopadd!(a_dest, op1, a1, op2, a2, α, β; kwargs...)

Compute `a_dest = α * op1(a1) * op2(a2) + β * a_dest`, matching dimensions by name. `op1` and
`op2` can be `identity` or `conj`, and keyword arguments (such as algorithm selection) are passed to `TensorAlgebra.contractopadd!`.
"""
function mulopadd!(
a_dest::AbstractNamedTensor,
op1, a1::AbstractNamedTensor,
op2, a2::AbstractNamedTensor,
α::Number, β::Number;
kwargs...
)
TA.contractopadd!(
unnamed(a_dest), names(a_dest),
op1, unnamed(a1), names(a1),
op2, unnamed(a2), names(a2),
α, β; kwargs...
)
return a_dest
end
Expand Down
2 changes: 1 addition & 1 deletion test/test_exports.jl
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ using Test: @test, @testset
:tryuniqueind, :uniqueind, :uniqueinds, :unioninds, :uniquename,
]
publics = [
:IndexName, :name, :names, :setname, :space, :unnamed,
:IndexName, :mulopadd!, :name, :names, :setname, :space, :unnamed,
:unnamedtype,
:decoration, :emptytags, :gettag, :gettags, :hastag, :plev, :settags, :tags,
:unsettags,
Expand Down
56 changes: 53 additions & 3 deletions test/test_tensoralgebra.jl
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
using ITensorBase: ITensorBase, Index, NamedOneTo, id, inds, name, names, operator, prime,
rename, unname, unnamed
using LinearAlgebra: norm, tr
using ITensorBase: ITensorBase, Index, NamedOneTo, id, inds, mulopadd!, name, names,
operator, prime, rename, unname, unnamed
using LinearAlgebra: mul!, norm, tr
using MatrixAlgebraKit: left_null, left_orth, left_polar, lq_compact, lq_full, qr_compact,
qr_full, right_null, right_orth, right_polar, svd_compact, svd_trunc, svd_vals
using StableRNGs: StableRNG
Expand Down Expand Up @@ -225,3 +225,53 @@ using Test: @test, @test_broken, @testset
@test name(rn) != name(i)
end
end

# Records each call, then contracts with the default algorithm.
struct RecordingContract <: TensorAlgebra.ContractAlgorithm
calls::Base.RefValue{Int}
end
function TensorAlgebra.contractpermopadd!(
alg::RecordingContract, a_dest, perm_dest_codomain, perm_dest_domain,
op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain,
α::Number, β::Number
)
alg.calls[] += 1
return TensorAlgebra.contractpermopadd!(
TensorAlgebra.MatricizeContract(), a_dest, perm_dest_codomain, perm_dest_domain,
op1, a1, perm1_codomain, perm1_domain, op2, a2, perm2_codomain, perm2_domain, α, β
)
end

@testset "`mul!` forwards `alg`" begin
i, j, k = NamedOneTo(2, "i"), NamedOneTo(3, "j"), NamedOneTo(4, "k")
a, b = randn(i, j), randn(j, k)
alg = RecordingContract(Ref(0))
dest = similar(a * b)
mul!(dest, a, b; alg)
@test dest ≈ a * b
@test alg.calls[] == 1
mul!(dest, a, b, 2, 1; alg)
@test dest ≈ 3 * (a * b)
@test alg.calls[] == 2
end

@testset "`mulopadd!`" begin
i, j, k = NamedOneTo(2, "i"), NamedOneTo(3, "j"), NamedOneTo(4, "k")
a, b = randn(ComplexF64, i, j), randn(ComplexF64, j, k)
dest = randn(ComplexF64, i, k)
mulopadd!(dest, conj, a, identity, b, true, false)
@test dest ≈ conj(a) * b
mulopadd!(dest, identity, a, conj, b, true, false)
@test dest ≈ a * conj(b)
prev = copy(dest)
mulopadd!(dest, conj, a, identity, b, 2, 1)
@test dest ≈ prev + 2 * (conj(a) * b)
alg = RecordingContract(Ref(0))
mulopadd!(dest, conj, a, identity, b, true, false; alg)
@test dest ≈ conj(a) * b
@test alg.calls[] == 1
dest_perm = randn(ComplexF64, k, i)
mulopadd!(dest_perm, conj, a, identity, b, true, false)
@test names(dest_perm) == ["k", "i"]
@test dest_perm ≈ conj(a) * b
end
Loading