diff --git a/Project.toml b/Project.toml index 7a5ada9..4955010 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "ITensorBase" uuid = "4795dd04-0d67-49bb-8f44-b89c448a1dc7" -version = "0.15.2" +version = "0.15.3" authors = ["ITensor developers and contributors"] [workspace] diff --git a/src/ITensorBase.jl b/src/ITensorBase.jl index 04471b3..983be9e 100644 --- a/src/ITensorBase.jl +++ b/src/ITensorBase.jl @@ -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 diff --git a/src/tensoralgebra.jl b/src/tensoralgebra.jl index d40c6c9..4ea92b9 100644 --- a/src/tensoralgebra.jl +++ b/src/tensoralgebra.jl @@ -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 diff --git a/test/test_exports.jl b/test/test_exports.jl index 0c529b4..dd608de 100644 --- a/test/test_exports.jl +++ b/test/test_exports.jl @@ -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, diff --git a/test/test_tensoralgebra.jl b/test/test_tensoralgebra.jl index 8130f69..207c38c 100644 --- a/test/test_tensoralgebra.jl +++ b/test/test_tensoralgebra.jl @@ -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 @@ -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