add rocshmem_barrier() (#61)

* add team-barrier implementation

add a team-barrier API and implementation in the IPC and RO conduit.
Clean up some of the logic in the RO Conduit to distinguish between
sync, sync_all, barrier, and barrier_all.

* add team_barrier_tests to functional tests
This commit is contained in:
Edgar Gabriel
2025-03-24 11:23:03 -05:00
کامیت شده توسط GitHub
والد e8ba20c5f5
کامیت bcbc42e78f
18فایلهای تغییر یافته به همراه271 افزوده شده و 14 حذف شده
@@ -39,7 +39,7 @@ enum ro_net_cmds {
RO_NET_FINALIZE,
RO_NET_TEAM_REDUCE,
RO_NET_SYNC,
RO_NET_BARRIER_ALL,
RO_NET_BARRIER,
RO_NET_TEAM_BROADCAST,
RO_NET_ALLTOALL,
RO_NET_FCOLLECT,
@@ -165,16 +165,26 @@ __device__ void *ROContext::shmem_ptr(const void *dest, int pe) {
__device__ void ROContext::barrier_all() {
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_BARRIER_ALL, nullptr, nullptr, 0, 0, 0, 0, 0,
build_queue_element(RO_NET_BARRIER, nullptr, nullptr, 0, 0, 0, 0, 0,
nullptr, nullptr, (MPI_Comm)NULL, ro_net_win_id,
block_handle, true, get_status_flag());
}
__syncthreads();
}
__device__ void ROContext::barrier(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_BARRIER, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
nullptr, team_obj->mpi_comm, ro_net_win_id, block_handle,
true, get_status_flag());
}
__syncthreads();
}
__device__ void ROContext::sync_all() {
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_BARRIER_ALL, nullptr, nullptr, 0, 0, 0, 0, 0,
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0,
nullptr, nullptr, (MPI_Comm)NULL, ro_net_win_id,
block_handle, true, get_status_flag());
}
@@ -65,6 +65,8 @@ class ROContext : public Context {
__device__ void barrier_all();
__device__ void barrier(rocshmem_team_t team);
__device__ void sync_all();
__device__ void sync(rocshmem_team_t team);
@@ -201,12 +201,16 @@ void MPITransport::submitRequestsToMPI() {
next_element.dst, next_element.src, next_element.ol1.size,
next_element.team_comm);
break;
case RO_NET_BARRIER_ALL:
barrier(queue_idx, next_element.status, true, ro_net_comm_world);
case RO_NET_BARRIER:
barrier(queue_idx, next_element.status, true,
next_element.team_comm == NULL ? ro_net_comm_world : next_element.team_comm,
true);
DPRINTF("Received Barrier_all\n");
break;
case RO_NET_SYNC:
barrier(queue_idx, next_element.status, true, next_element.team_comm);
barrier(queue_idx, next_element.status, true,
next_element.team_comm == NULL ? ro_net_comm_world : next_element.team_comm,
false);
DPRINTF("Received Sync\n");
break;
case RO_NET_FENCE:
@@ -269,12 +273,19 @@ void MPITransport::global_exit(int status) {
}
void MPITransport::barrier(int contextId, volatile char *status, bool blocking,
MPI_Comm team) {
MPI_Comm team, bool do_quiet) {
MPI_Request request{};
NET_CHECK(MPI_Ibarrier(team, &request));
requests.push_back({request, {status, contextId, blocking}});
outstanding[contextId]++;
if (do_quiet) {
requests.push_back({request, {nullptr, contextId, false}});
outstanding[contextId]++;
quiet(contextId, status);
} else {
requests.push_back({request, {status, contextId, blocking}});
outstanding[contextId]++;
}
}
MPI_Op MPITransport::get_mpi_op(ROCSHMEM_OP op) {
@@ -388,7 +399,7 @@ void MPITransport::team_broadcast(void *dst, void *src, int size, int win_id,
}
NET_CHECK(MPI_Win_flush_all(bp->heap_window_info[win_id]->get_win()));
barrier(contextId, nullptr, false, comm);
barrier(contextId, nullptr, false, comm, false);
quiet(contextId, status);
}
@@ -52,7 +52,7 @@ public:
rocshmem_team_t *new_team) override;
void barrier(int contextId, volatile char *status, bool blocking,
MPI_Comm team) override;
MPI_Comm team, bool quiet) override;
void team_reduction(void *dst, void *src, int size, int win_id,
int contextId, MPI_Comm team, ROCSHMEM_OP op,
@@ -51,7 +51,7 @@ class Transport {
rocshmem_team_t *new_team) = 0;
virtual void barrier(int wg_id, volatile char *status, bool blocking,
MPI_Comm team) = 0;
MPI_Comm team, bool quiet) = 0;
virtual void team_reduction(void *dst, void *src, int size, int win_id,
int wg_id, MPI_Comm team, ROCSHMEM_OP op,