Update Barrier and Sync APIs (#73)

* Add thread, wavefront, and workgroup-level `barrier` APIs in IPC and RO conduits; remove collectives on default context
 - Implemented `barrier` APIs for thread, wavefront, and workgroup scopes
 - Added support into both IPC and RO conduits
 - Added functional tests to cover all `barrier` APIs
 - Removed collective operations on default context

* Add thread, wavefront, and workgroup-level `sync` APIs in IPC and RO conduits.
  - Implemented `sync` APIs for thread, wavefront, and workgroup scopes
  - Added support into both IPC and RO conduits
  - Added functional tests to cover all `sync` APIs

* update naming convention for context-based `barrier` APIs
Αυτή η υποβολή περιλαμβάνεται σε:
Avinash Kethineedi
2025-04-08 11:25:31 -05:00
υποβλήθηκε από GitHub
γονέας c652f58cef
υποβολή dc61bca066
16 αρχεία άλλαξαν με 347 προσθήκες και 67 διαγραφές
@@ -193,6 +193,22 @@ __device__ void ROContext::barrier_all_wg() {
}
__device__ void ROContext::barrier(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
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(), is_default_ctx);
}
__device__ void ROContext::barrier_wave(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
if (is_thread_zero_in_wave()) {
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(), is_default_ctx);
}
}
__device__ void ROContext::barrier_wg(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,
@@ -225,6 +241,22 @@ __device__ void ROContext::sync_all_wg() {
__syncthreads();
}
__device__ void ROContext::sync(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
nullptr, team_obj->mpi_comm, ro_net_win_id, block_handle,
true, get_status_flag(), is_default_ctx);
}
__device__ void ROContext::sync_wave(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
if (is_thread_zero_in_wave()) {
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
nullptr, team_obj->mpi_comm, ro_net_win_id, block_handle,
true, get_status_flag(), is_default_ctx);
}
}
__device__ void ROContext::sync_wg(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
if (is_thread_zero_in_block()) {
@@ -71,12 +71,20 @@ class ROContext : public Context {
__device__ void barrier(rocshmem_team_t team);
__device__ void barrier_wave(rocshmem_team_t team);
__device__ void barrier_wg(rocshmem_team_t team);
__device__ void sync_all();
__device__ void sync_all_wave();
__device__ void sync_all_wg();
__device__ void sync(rocshmem_team_t team);
__device__ void sync_wave(rocshmem_team_t team);
__device__ void sync_wg(rocshmem_team_t team);
template <typename T>