Remove MPI compile-time dependency (#264)

* use dlsym for MPI functions

to allow compiling without MPI support, convert the usage of MPI functions and symbols to be based on a dlopen/dlsym based mechanism. Turns out this cannot be done entirely vendor neutral, slightly different solutions might be required for Open MPI, MPICH and the new MPI ABI.

* checkpoint

more work to be done.

* checkpoint 2

* checkpoint 3

* checkpoint 4

examples compile and link correctly

* checkpoitn 5 (I think)

* Checkpoitn 6

* dyld-mpi: adapt GDA

* dyldmpi: tests that depend on MPI need to link with it themselves

* do not ../mpi_instance.h

* dyldmpi: make the symetricHeapTestFixture compile

* dyldmpi: Change cmakery, compiles and run gda w/o external MPI

* Make it also compile in external MPI mode

* dyldmpi: ipc unit tests compile but do not link

* dyldmpi: new approach, if external mpi required, link with mpi,
otherwise use ompi5 abi

* C-style comments in cmakelist..

* dyldmpi: examples: do not fail compiling if MPI not found at build time,
instead do not compile the MPI required examples

* more updates to CMake logic

* convert RO backend

and a few other cleanups

* update some unit tests

to work with the dlopen MPI environment correctly.

---------

Co-authored-by: Aurelien Bouteiller <abouteil@amd.com>

[ROCm/rocshmem commit: e4c427a736]
This commit is contained in:
Edgar Gabriel
2025-10-01 08:06:56 -05:00
committed by GitHub
parent 4f955324ac
commit 53fa35b980
50 changed files with 712 additions and 294 deletions
@@ -72,8 +72,9 @@ target_sources(
# ROCSHMEM
###############################################################################
if (BUILD_TESTS_ONLY)
find_package(MPI REQUIRED)
find_package(hip REQUIRED PATHS /opt/rocm)
#TODO check that build_test_only still works with external-mpi
#find_package(MPI REQUIRED)
#find_package(hip REQUIRED PATHS /opt/rocm)
find_package(rocshmem REQUIRED PATHS /opt/rocm)
target_include_directories(
@@ -25,7 +25,6 @@
#include "tester.hpp"
#include <hip/hip_runtime.h>
#include <mpi.h>
#include <functional>
#include <iostream>
@@ -79,9 +79,9 @@ endif()
# ROCSHMEM DEPENDENCY
###############################################################################
find_package(hip REQUIRED PATHS /opt/rocm)
find_package(MPI REQUIRED)
if (BUILD_TESTS_ONLY)
find_package(MPI REQUIRED)
find_package(rocshmem REQUIRED PATHS /opt/rocm)
target_include_directories(
@@ -95,6 +95,7 @@ endif()
target_link_libraries(
${PROJECT_NAME}
PRIVATE
MPI::MPI_CXX
roc::rocshmem
)
@@ -32,13 +32,13 @@ TEST_P(DegenerateSimpleCoarse, ptr_check) {
}
TEST_P(DegenerateSimpleCoarse, MPI_num_pes) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
}
TEST_P(DegenerateSimpleCoarse, IPC_bases) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
ASSERT_NE(ipc_impl_.ipc_bases, nullptr);
for(int i{0}; i < mpi_.num_pes(); i++) {
for(int i{0}; i < mpi_->num_pes(); i++) {
ASSERT_NE(ipc_impl_.ipc_bases[i], nullptr);
}
}
@@ -73,7 +73,10 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
public:
IPCImplSimpleCoarse() {
ipc_impl_.ipcHostInit(mpi_.my_pe(), mpi_.get_heap_bases() , MPI_COMM_WORLD);
MPIInstance::mpilib_dl_init();
mpi_ = new MPI_T (heap_mem_.get_ptr(), heap_mem_.get_size(), MPI_COMM_WORLD);
ipc_impl_.ipcHostInit(mpi_->my_pe(), mpi_->get_heap_bases(), MPI_COMM_WORLD);
assert(ipc_impl_dptr_ == nullptr);
hip_allocator_.allocate((void**)&ipc_impl_dptr_, sizeof(IpcImpl));
CHECK_HIP(hipMemcpy(ipc_impl_dptr_, &ipc_impl_,
@@ -85,6 +88,7 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
hip_allocator_.deallocate(ipc_impl_dptr_);
}
ipc_impl_.ipcHostStop();
MPIInstance::mpilib_dl_close();
}
void launch(FN_T f, const dim3 grid, const dim3 block, int* src, int* dest, size_t bytes) {
@@ -132,7 +136,7 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
return;
}
size_t bytes = golden_.size() * sizeof(int);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
CHECK_HIP(hipMemcpy(dev_src, golden_.data(), bytes, hipMemcpyHostToDevice));
CHECK_HIP(hipStreamSynchronize(nullptr));
}
@@ -140,14 +144,14 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
bool pe_initializes_src_buffer(TestType test) {
bool is_write_test = test;
bool is_read_test = !test;
return (is_write_test && mpi_.my_pe() == 0) ||
(is_read_test && mpi_.my_pe() == 1);
return (is_write_test && mpi_->my_pe() == 0) ||
(is_read_test && mpi_->my_pe() == 1);
}
void execute(TestType test, FN_T fn, const dim3 grid, const dim3 block) {
if (mpi_.my_pe()) {
mpi_.barrier();
mpi_.barrier();
if (mpi_->my_pe()) {
mpi_->barrier();
mpi_->barrier();
return;
}
int *src{nullptr};
@@ -160,9 +164,9 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[0]);
}
size_t bytes = golden_.size() * sizeof(int);
mpi_.barrier();
mpi_->barrier();
launch(fn, grid, block, src, dest, bytes);
mpi_.barrier();
mpi_->barrier();
}
void validate_dest_buffer(TestType test) {
@@ -170,7 +174,7 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
return;
}
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
for (int i = 0; i < static_cast<int>(golden_.size()); i++) {
ASSERT_EQ(golden_[i], dev_dest[i]);
}
@@ -184,7 +188,7 @@ class IPCImplSimpleCoarse : public ::testing::TestWithParam<std::tuple<int, int,
std::vector<int> golden_;
HEAP_T heap_mem_ {};
MPI_T mpi_ {heap_mem_.get_ptr(), heap_mem_.get_size()};
MPI_T *mpi_{nullptr};
IpcImpl ipc_impl_ {};
IpcImpl *ipc_impl_dptr_ {nullptr};
@@ -31,13 +31,13 @@ TEST_P(DegenerateSimpleFine, ptr_check) {
}
TEST_P(DegenerateSimpleFine, MPI_num_pes) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
}
TEST_P(DegenerateSimpleFine, IPC_bases) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
ASSERT_NE(ipc_impl_.ipc_bases, nullptr);
for(int i{0}; i < mpi_.num_pes(); i++) {
for(int i{0}; i < mpi_->num_pes(); i++) {
ASSERT_NE(ipc_impl_.ipc_bases[i], nullptr);
}
}
@@ -140,7 +140,10 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
public:
IPCImplSimpleFine() {
ipc_impl_.ipcHostInit(mpi_.my_pe(), mpi_.get_heap_bases() , MPI_COMM_WORLD);
MPIInstance::mpilib_dl_init();
mpi_ = new MPI_T (heap_mem_.get_ptr(), heap_mem_.get_size(), MPI_COMM_WORLD);
ipc_impl_.ipcHostInit(mpi_->my_pe(), mpi_->get_heap_bases(), MPI_COMM_WORLD);
assert(ipc_impl_dptr_ == nullptr);
hip_allocator_.allocate((void**)&ipc_impl_dptr_, sizeof(IpcImpl));
@@ -163,6 +166,7 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
}
ipc_impl_.ipcHostStop();
MPIInstance::mpilib_dl_close();
}
void launch(FN_T1 f, const dim3 grid, const dim3 block, int* src, int* dest, size_t bytes, TestType test) {
@@ -214,7 +218,7 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
void initialize_signal(TestType test) {
bool is_write_test = test;
if (is_write_test && mpi_.my_pe() == 0) {
if (is_write_test && mpi_->my_pe() == 0) {
int *dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
*(dest + SIGNAL_OFFSET) = 0;
}
@@ -225,27 +229,27 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
return;
}
size_t bytes = golden_.size() * sizeof(int);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
CHECK_HIP(hipMemcpy(dev_src, golden_.data(), bytes, hipMemcpyHostToDevice));
}
bool pe_initializes_src_buffer(TestType test) {
bool is_write_test = test;
bool is_read_test = !test;
return (is_write_test && mpi_.my_pe() == 0) ||
(is_read_test && mpi_.my_pe() == 1);
return (is_write_test && mpi_->my_pe() == 0) ||
(is_read_test && mpi_->my_pe() == 1);
}
void execute(TestType test, FN_T1 fn, const dim3 grid, const dim3 block) {
size_t bytes = golden_.size() * sizeof(int);
if (mpi_.my_pe()) {
mpi_.barrier();
if (mpi_->my_pe()) {
mpi_->barrier();
if (test == WRITE) {
int *dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
FN_T2 val_fn = kernel_put_with_signal_simple_validator;
launch(val_fn, grid, block, dest, bytes);
}
mpi_.barrier();
mpi_->barrier();
return;
}
int *src{nullptr};
@@ -257,9 +261,9 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[0]);
}
mpi_.barrier();
mpi_->barrier();
launch(fn, grid, block, src, dest, bytes, test);
mpi_.barrier();
mpi_->barrier();
}
void check_device_validation_errors(TestType test) {
@@ -274,7 +278,7 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
return;
}
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
for (int i = 0; i < static_cast<int>(golden_.size()); i++) {
ASSERT_EQ(golden_[i], dev_dest[i]);
}
@@ -291,7 +295,7 @@ class IPCImplSimpleFine : public ::testing::TestWithParam<std::tuple<int, int, i
HEAP_T heap_mem_ {};
MPI_T mpi_ {heap_mem_.get_ptr(), heap_mem_.get_size()};
MPI_T *mpi_{nullptr};
std::vector<int> golden_;
@@ -33,13 +33,13 @@ TEST_F(DegenerateTiledFine, ptr_check) {
}
TEST_F(DegenerateTiledFine, MPI_num_pes) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
}
TEST_F(DegenerateTiledFine, IPC_bases) {
ASSERT_EQ(mpi_.num_pes(), 2);
ASSERT_EQ(mpi_->num_pes(), 2);
ASSERT_NE(ipc_impl_.ipc_bases, nullptr);
for(int i{0}; i < mpi_.num_pes(); i++) {
for(int i{0}; i < mpi_->num_pes(); i++) {
ASSERT_NE(ipc_impl_.ipc_bases[i], nullptr);
}
}
@@ -158,7 +158,10 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
public:
IPCImplTiledFine() {
ipc_impl_.ipcHostInit(mpi_.my_pe(), mpi_.get_heap_bases() , MPI_COMM_WORLD);
MPIInstance::mpilib_dl_init();
mpi_ = new MPI_T (heap_mem_.get_ptr(), heap_mem_.get_size(), MPI_COMM_WORLD);
ipc_impl_.ipcHostInit(mpi_->my_pe(), mpi_->get_heap_bases(), MPI_COMM_WORLD);
assert(ipc_impl_dptr_ == nullptr);
hip_allocator_.allocate((void**)&ipc_impl_dptr_, sizeof(IpcImpl));
@@ -181,6 +184,7 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
}
ipc_impl_.ipcHostStop();
MPIInstance::mpilib_dl_close();
}
void launch(FN_T1 f, const dim3 grid, const dim3 block, int* src, int* dest, size_t bytes, TestType test) {
@@ -232,7 +236,7 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
void initialize_signal(TestType test, int signal_value = 0) {
bool is_write_test = test;
if (is_write_test && mpi_.my_pe() == 0) {
if (is_write_test && mpi_->my_pe() == 0) {
int *dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
*(dest + SIGNAL_OFFSET) = signal_value;
}
@@ -243,28 +247,28 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
return;
}
size_t bytes = golden_.size() * sizeof(int);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
CHECK_HIP(hipMemcpy(dev_src, golden_.data(), bytes, hipMemcpyHostToDevice));
}
bool pe_initializes_src_buffer(TestType test) {
bool is_write_test = test;
bool is_read_test = !test;
return (is_write_test && mpi_.my_pe() == 0) ||
(is_read_test && mpi_.my_pe() == 1);
return (is_write_test && mpi_->my_pe() == 0) ||
(is_read_test && mpi_->my_pe() == 1);
}
void execute(TestType test, FN_T1 fn, const dim3 grid, const dim3 block) {
size_t bytes = golden_.size() * sizeof(int);
if (mpi_.my_pe()) {
mpi_.barrier();
if (mpi_->my_pe()) {
mpi_->barrier();
if (test == WRITE) {
int *dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
FN_T2 val_fn = kernel_put_with_signal_tiled_validator;
launch(val_fn, grid, block, dest, bytes);
ASSERT_EQ(*(dest + SIGNAL_OFFSET), 0);
}
mpi_.barrier();
mpi_->barrier();
return;
}
int *src{nullptr};
@@ -276,9 +280,9 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
src = reinterpret_cast<int*>(ipc_impl_.ipc_bases[1]);
dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[0]);
}
mpi_.barrier();
mpi_->barrier();
launch(fn, grid, block, src, dest, bytes, test);
mpi_.barrier();
mpi_->barrier();
}
void check_device_validation_errors(TestType test) {
@@ -293,7 +297,7 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
return;
}
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_.my_pe()]);
auto dev_dest = reinterpret_cast<int*>(ipc_impl_.ipc_bases[mpi_->my_pe()]);
for (int i = 0; i < static_cast<int>(golden_.size()); i++) {
ASSERT_EQ(golden_[i], dev_dest[i]);
}
@@ -310,7 +314,7 @@ class IPCImplTiledFine : public ::testing::TestWithParam<std::tuple<int, int, in
HEAP_T heap_mem_ {};
MPI_T mpi_ {heap_mem_.get_ptr(), heap_mem_.get_size()};
MPI_T *mpi_ {nullptr};
std::vector<int> golden_;
@@ -27,6 +27,8 @@
#include "gtest/gtest.h"
#include <mpi.h>
#include "../src/memory/heap_memory.hpp"
#include "../src/memory/hip_allocator.hpp"
#include "../src/memory/remote_heap_info.hpp"
@@ -55,7 +57,8 @@ class RemoteHeapInfoTestFixture : public ::testing::Test
* @brief Remote heap info with MPI Communicator
*/
MPI_T mpi_ {heap_mem_.get_ptr(),
heap_mem_.get_size()};
heap_mem_.get_size(),
MPI_COMM_WORLD};
};
} // namespace rocshmem
@@ -30,27 +30,27 @@ TEST_F(SymmetricHeapTestFixture, malloc_free) {
void *ptr{nullptr};
size_t request_bytes{48};
symmetric_heap_.malloc(&ptr, request_bytes);
symmetric_heap_->malloc(&ptr, request_bytes);
ASSERT_NE(ptr, nullptr);
ASSERT_NO_FATAL_FAILURE(symmetric_heap_.free(ptr));
ASSERT_NO_FATAL_FAILURE(symmetric_heap_->free(ptr));
}
TEST_F(SymmetricHeapTestFixture, window_info) {
auto win_info_ptr{symmetric_heap_.get_window_info()};
auto win_info_ptr{symmetric_heap_->get_window_info()};
WindowInfoMPI* window_info_mpi = dynamic_cast<WindowInfoMPI*>(win_info_ptr);
if (window_info_mpi) {
void *window_base_addr{nullptr};
int flag{0};
MPI_Win_get_attr(window_info_mpi->get_win(), MPI_WIN_BASE, &window_base_addr,
&flag);
&flag);
ASSERT_NE(0, flag);
ASSERT_NE(nullptr, window_base_addr);
}
}
TEST_F(SymmetricHeapTestFixture, heap_bases) {
auto heap_bases{symmetric_heap_.get_heap_bases()};
auto heap_bases{symmetric_heap_->get_heap_bases()};
for (const auto &base : heap_bases) {
ASSERT_NE(nullptr, base);
}
@@ -25,6 +25,8 @@
#ifndef ROCSHMEM_SYMMETRIC_HEAP_GTEST_HPP
#define ROCSHMEM_SYMMETRIC_HEAP_GTEST_HPP
#include <mpi.h>
#include "gtest/gtest.h"
#include "../src/memory/symmetric_heap.hpp"
@@ -37,7 +39,16 @@ class SymmetricHeapTestFixture : public ::testing::Test
/**
* @brief Symmetric heap object
*/
SymmetricHeap symmetric_heap_ {MPI_COMM_WORLD};
SymmetricHeap *symmetric_heap_;
void SetUp() override {
MPIInstance::mpilib_dl_init();
symmetric_heap_ = new SymmetricHeap(MPI_COMM_WORLD);
}
void TearDown() override {
MPIInstance::mpilib_dl_close();
}
};
} // namespace rocshmem