diff --git a/Project.toml b/Project.toml index dfafa378..9d36bded 100644 --- a/Project.toml +++ b/Project.toml @@ -26,6 +26,7 @@ WrappedUnions = "325db55a-9c6c-5b90-b1a2-ec87e7a38c44" [weakdeps] Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e" GradedArrays = "bc96ca6e-b7c8-4bb6-888e-c93f838762c2" +MPI = "da04e1cc-30fd-572f-bb4f-1f8673147195" Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" OMEinsumContractionOrders = "6f22d1fd-8eed-4bb7-9776-e7d684900715" TensorKit = "07d1fe3e-3e46-537d-9eac-e9e13d0d4cec" @@ -34,6 +35,7 @@ TensorKitSectors = "13a9c161-d5da-41f0-bcbd-e1a08ae0647f" [extensions] ITensorBaseAdaptExt = "Adapt" ITensorBaseGradedArraysExt = ["GradedArrays", "TensorKitSectors"] +ITensorBaseMPIExt = "MPI" ITensorBaseMooncakeExt = "Mooncake" ITensorBaseOMEinsumContractionOrdersExt = "OMEinsumContractionOrders" ITensorBaseTensorKitExt = "TensorKit" @@ -47,6 +49,7 @@ Combinatorics = "1" ConstructionBase = "1.6" GradedArrays = "0.16" LinearAlgebra = "1.10" +MPI = "0.20.26" MatrixAlgebraKit = "0.2, 0.3, 0.4, 0.5, 0.6" Mooncake = "0.4.202, 0.5" OMEinsumContractionOrders = "1.3" diff --git a/ext/ITensorBaseMPIExt.jl b/ext/ITensorBaseMPIExt.jl new file mode 100644 index 00000000..60726459 --- /dev/null +++ b/ext/ITensorBaseMPIExt.jl @@ -0,0 +1,9 @@ +module ITensorBaseMPIExt + +using ITensorBase: AbstractNamedArray, AbstractNamedTensor, unnamed +using MPI: MPI + +MPI.Buffer(a::AbstractNamedArray) = MPI.Buffer(unnamed(a)) +MPI.Buffer(a::AbstractNamedTensor) = MPI.Buffer(unnamed(a)) + +end diff --git a/test/Project.toml b/test/Project.toml index 7e3740f5..88276e52 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -8,6 +8,7 @@ ITensorBase = "4795dd04-0d67-49bb-8f44-b89c448a1dc7" ITensorPkgSkeleton = "3d388ab1-018a-49f4-ae50-18094d5f71ea" JLArrays = "27aeb0d3-9eb9-45fb-866b-73c2ecf80fcb" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" +MPI = "da04e1cc-30fd-572f-bb4f-1f8673147195" MatrixAlgebraKit = "6c742aac-3347-4629-af66-fc926824e5e4" Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" OMEinsumContractionOrders = "6f22d1fd-8eed-4bb7-9776-e7d684900715" @@ -37,6 +38,7 @@ ITensorBase = "0.14" ITensorPkgSkeleton = "0.3.42" JLArrays = "0.2, 0.3" LinearAlgebra = "1.10" +MPI = "0.20.26" MatrixAlgebraKit = "0.2, 0.3, 0.4, 0.5, 0.6" Mooncake = "0.4, 0.5" OMEinsumContractionOrders = "1.3" diff --git a/test/test_mpiext.jl b/test/test_mpiext.jl new file mode 100644 index 00000000..4420fb78 --- /dev/null +++ b/test/test_mpiext.jl @@ -0,0 +1,31 @@ +using ITensorBase: NamedArray, nameddims, unnamed +using MPI: MPI +using Test: @test, @testset + +@testset "MPIExt (eltype=$elt)" for elt in (Float64, ComplexF64) + @testset "Buffer wraps the unnamed parent" begin + nt = nameddims(randn(elt, (2, 3)), ("i", "j")) + na = NamedArray(randn(elt, 4), "x") + for a in (nt, na) + buffer = MPI.Buffer(a) + @test buffer.data ≡ unnamed(a) + @test buffer.count == length(unnamed(a)) + @test buffer.datatype == MPI.Datatype(elt) + end + end + @testset "Sendrecv! round trip" begin + MPI.Initialized() || MPI.Init() + comm = MPI.COMM_WORLD + rank = MPI.Comm_rank(comm) + + send = nameddims(randn(elt, (2, 3)), ("i", "j")) + recv = nameddims(zeros(elt, (2, 3)), ("i", "j")) + MPI.Sendrecv!(send, recv, comm; dest = rank, source = rank) + @test unnamed(recv) == unnamed(send) + + send = NamedArray(randn(elt, 4), "x") + recv = NamedArray(zeros(elt, 4), "x") + MPI.Sendrecv!(send, recv, comm; dest = rank, source = rank) + @test unnamed(recv) == unnamed(send) + end +end