Update Barrier_All and Sync_All APIs (#72)

* Fix deadlock in `rocshmem_ctx_wg_barrier_all` API in IPC conduit by adding per-context pSync buffers and context IDs
  - Added separate pSync buffers for each device context
  - Resolved deadlock when invoking barrier API (`rocshmem_ctx_wg_barrier_all`) concurrently from multiple contexts

* Update barrier_all functional tests for multi-context support

* Add thread, wavefront, and workgroup-level barrier_all APIs in IPC and RO conduits
  - Implemented barrier_all APIs at thread, wavefront, and workgroup granularity
  - Added support in both IPC and RO conduits
  - Updated functional tests to cover all `barrier_all` APIs

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

[ROCm/rocshmem commit: c652f58cef]
Этот коммит содержится в:
Avinash Kethineedi
2025-04-02 11:58:55 -05:00
коммит произвёл GitHub
родитель 0cde5f53dc
Коммит 426bbf525b
22 изменённых файлов: 508 добавлений и 53 удалений
+4
Просмотреть файл
@@ -130,6 +130,8 @@ void Backend::dump_stats() {
printf("Quiets %llu\n", device_stats.getStat(NUM_QUIET));
printf("ToAll %llu\n", device_stats.getStat(NUM_TO_ALL));
printf("BarrierAll %llu\n", device_stats.getStat(NUM_BARRIER_ALL));
printf("WAVE_BarrierAll %llu\n", device_stats.getStat(NUM_BARRIER_ALL_WAVE));
printf("WG_BarrierAll %llu\n", device_stats.getStat(NUM_BARRIER_ALL_WG));
printf("Wait Until %llu\n", device_stats.getStat(NUM_WAIT_UNTIL));
printf("Wait Until Any %llu\n", device_stats.getStat(NUM_WAIT_UNTIL_ANY));
printf("Wait Until All %llu\n", device_stats.getStat(NUM_WAIT_UNTIL_ALL));
@@ -153,6 +155,8 @@ void Backend::dump_stats() {
printf("Tests %llu\n", device_stats.getStat(NUM_TEST));
printf("SHMEM_PTR %llu\n", device_stats.getStat(NUM_SHMEM_PTR));
printf("SyncAll %llu\n", device_stats.getStat(NUM_SYNC_ALL));
printf("WAVE_SyncAll %llu\n", device_stats.getStat(NUM_SYNC_ALL_WAVE));
printf("WG_SyncAll %llu\n", device_stats.getStat(NUM_SYNC_ALL_WG));
const auto& host_stats{globalHostStats};
printf("HOST STATS\n");
+9 -1
Просмотреть файл
@@ -137,11 +137,19 @@ class Context {
__device__ void barrier_all();
__device__ void barrier_all_wave();
__device__ void barrier_all_wg();
__device__ void barrier(rocshmem_team_t team);
__device__ void sync_all();
__device__ void sync(rocshmem_team_t team);
__device__ void sync_all_wave();
__device__ void sync_all_wg();
__device__ void sync_wg(rocshmem_team_t team);
template <typename T>
__device__ T amo_fetch(void* dst, T value, T cond, int pe, uint8_t atomic_op);
+27 -3
Просмотреть файл
@@ -148,6 +148,18 @@ __device__ void Context::barrier_all() {
DISPATCH(barrier_all());
}
__device__ void Context::barrier_all_wave() {
ctxStats.incStat(NUM_BARRIER_ALL_WAVE);
DISPATCH(barrier_all_wave());
}
__device__ void Context::barrier_all_wg() {
ctxStats.incStat(NUM_BARRIER_ALL_WG);
DISPATCH(barrier_all_wg());
}
__device__ void Context::barrier(rocshmem_team_t team) {
ctxStats.incStat(NUM_BARRIER_ALL);
@@ -160,10 +172,22 @@ __device__ void Context::sync_all() {
DISPATCH(sync_all());
}
__device__ void Context::sync(rocshmem_team_t team) {
ctxStats.incStat(NUM_SYNC_ALL);
__device__ void Context::sync_all_wave() {
ctxStats.incStat(NUM_SYNC_ALL_WAVE);
DISPATCH(sync(team));
DISPATCH(sync_all_wave());
}
__device__ void Context::sync_all_wg() {
ctxStats.incStat(NUM_SYNC_ALL_WG);
DISPATCH(sync_all_wg());
}
__device__ void Context::sync_wg(rocshmem_team_t team) {
ctxStats.incStat(NUM_SYNC_ALL_WG);
DISPATCH(sync_wg(team));
}
__device__ void Context::putmem_wg(void* dest, const void* source,
+11 -6
Просмотреть файл
@@ -114,8 +114,9 @@ IPCBackend::~IPCBackend() {
void IPCBackend::setup_ctxs() {
CHECK_HIP(hipMalloc(&ctx_array, sizeof(IPCContext) * maximum_num_contexts_));
// 0th context is default context
for (size_t i = 0; i < maximum_num_contexts_; i++) {
new (&ctx_array[i]) IPCContext(this);
new (&ctx_array[i]) IPCContext(this, i + 1);
ctx_free_list.get()->push_back(ctx_array + i);
}
}
@@ -278,9 +279,10 @@ void IPCBackend::init_wrk_sync_buffer() {
auto max_num_teams{team_tracker.get_max_num_teams()};
/**
* size of barrier sync
* size of barrier sync for all the contexts
*/
Wrk_Sync_buffer_size_ += sizeof(*barrier_sync) * ROCSHMEM_BARRIER_SYNC_SIZE;
Wrk_Sync_buffer_size_ += sizeof(*barrier_sync) * ROCSHMEM_BARRIER_SYNC_SIZE *
(maximum_num_contexts_ + 1);
/**
* Size of sync arrays for the teams
@@ -378,15 +380,18 @@ void IPCBackend::rocshmem_collective_init() {
/*
* Allocate heap space for barrier_sync
*/
size_t one_sync_size_bytes{sizeof(*barrier_sync)};
size_t sync_size_bytes{one_sync_size_bytes * ROCSHMEM_BARRIER_SYNC_SIZE};
size_t one_sync_size_bytes {sizeof(*barrier_sync)};
size_t total_sync_elems {
ROCSHMEM_BARRIER_SYNC_SIZE * (maximum_num_contexts_ + 1)};
size_t sync_size_bytes {one_sync_size_bytes * total_sync_elems};
barrier_sync = reinterpret_cast<int64_t*>(temp_Wrk_Sync_buff_ptr_);
temp_Wrk_Sync_buff_ptr_ += sync_size_bytes;
/*
* Initialize the barrier synchronization array with default values.
*/
for (int i = 0; i < num_pes; i++) {
for (int i = 0; i < total_sync_elems; i++) {
barrier_sync[i] = ROCSHMEM_SYNC_VALUE;
}
+5 -2
Просмотреть файл
@@ -36,15 +36,18 @@
namespace rocshmem {
__host__ IPCContext::IPCContext(Backend *b)
__host__ IPCContext::IPCContext(Backend *b, unsigned int ctx_id)
: Context(b, false) {
IPCBackend *backend{static_cast<IPCBackend *>(b)};
ipcImpl_.ipc_bases = b->ipcImpl.ipc_bases;
ipcImpl_.shm_size = b->ipcImpl.shm_size;
barrier_sync = backend->barrier_sync;
size_t barrier_sync_offset = ctx_id * ROCSHMEM_BARRIER_SYNC_SIZE;
barrier_sync = backend->barrier_sync + barrier_sync_offset;
fence_pool = backend->fence_pool;
Wrk_Sync_buffer_bases_ = backend->get_wrk_sync_bases();
ctx_id_ = ctx_id;
orders_.store = detail::atomic::rocshmem_memory_order::memory_order_seq_cst;
}
+22 -3
Просмотреть файл
@@ -31,9 +31,9 @@ namespace rocshmem {
class IPCContext : public Context {
public:
__host__ IPCContext(Backend *b);
__host__ IPCContext(Backend *b, unsigned int ctx_id);
__device__ IPCContext(Backend *b);
__device__ IPCContext(Backend *b, unsigned int ctx_id);
__device__ void threadfence_system();
@@ -61,11 +61,19 @@ class IPCContext : public Context {
__device__ void barrier_all();
__device__ void barrier_all_wave();
__device__ void barrier_all_wg();
__device__ void barrier(rocshmem_team_t team);
__device__ void sync_all();
__device__ void sync(rocshmem_team_t team);
__device__ void sync_all_wave();
__device__ void sync_all_wg();
__device__ void sync_wg(rocshmem_team_t team);
template <typename T>
__device__ void p(T *dest, T value, int pe);
@@ -240,6 +248,12 @@ class IPCContext : public Context {
__device__ void internal_sync(int pe, int PE_start, int stride, int PE_size,
int64_t *pSync);
__device__ void internal_sync_wave(int pe, int PE_start, int stride, int PE_size,
int64_t *pSync);
__device__ void internal_sync_wg(int pe, int PE_start, int stride, int PE_size,
int64_t *pSync);
__device__ void internal_direct_barrier(int pe, int PE_start, int stride,
int n_pes, int64_t *pSync);
@@ -289,6 +303,11 @@ class IPCContext : public Context {
*/
char **Wrk_Sync_buffer_bases_{nullptr};
/**
* @brief Decive context Id
*/
unsigned int ctx_id_{};
public:
//TODO(Avinash):
//Make tinfo private variable, it requires changes to the context
+45 -5
Просмотреть файл
@@ -84,9 +84,29 @@ __device__ void IPCContext::internal_atomic_barrier(int pe, int PE_start,
}
}
// Uses PE values that are relative to world
__device__ void IPCContext::internal_sync(int pe, int PE_start, int stride,
int PE_size, int64_t *pSync) {
if (PE_size < 64) {
internal_direct_barrier(pe, PE_start, stride, PE_size, pSync);
} else {
internal_atomic_barrier(pe, PE_start, stride, PE_size, pSync);
}
}
__device__ void IPCContext::internal_sync_wave(int pe, int PE_start, int stride,
int PE_size, int64_t *pSync) {
if (is_thread_zero_in_wave()) {
if (PE_size < 64) {
internal_direct_barrier(pe, PE_start, stride, PE_size, pSync);
} else {
internal_atomic_barrier(pe, PE_start, stride, PE_size, pSync);
}
}
}
// Uses PE values that are relative to world
__device__ void IPCContext::internal_sync_wg(int pe, int PE_start, int stride,
int PE_size, int64_t *pSync) {
__syncthreads();
if (is_thread_zero_in_block()) {
if (PE_size < 64) {
@@ -98,7 +118,7 @@ __device__ void IPCContext::internal_sync(int pe, int PE_start, int stride,
__syncthreads();
}
__device__ void IPCContext::sync(rocshmem_team_t team) {
__device__ void IPCContext::sync_wg(rocshmem_team_t team) {
IPCTeam *team_obj = reinterpret_cast<IPCTeam *>(team);
int pe = team_obj->my_pe_in_world;
@@ -107,18 +127,38 @@ __device__ void IPCContext::sync(rocshmem_team_t team) {
int pe_size = team_obj->num_pes;
long *p_sync = team_obj->barrier_pSync;
internal_sync(pe, pe_start, pe_stride, pe_size, p_sync);
internal_sync_wg(pe, pe_start, pe_stride, pe_size, p_sync);
}
__device__ void IPCContext::sync_all() {
internal_sync(my_pe, 0, 1, num_pes, barrier_sync);
}
__device__ void IPCContext::sync_all_wave() {
internal_sync_wave(my_pe, 0, 1, num_pes, barrier_sync);
}
__device__ void IPCContext::sync_all_wg() {
internal_sync_wg(my_pe, 0, 1, num_pes, barrier_sync);
}
__device__ void IPCContext::barrier_all() {
quiet();
sync_all();
}
__device__ void IPCContext::barrier_all_wave() {
if (is_thread_zero_in_wave()) {
quiet();
}
sync_all_wave();
}
__device__ void IPCContext::barrier_all_wg() {
if (is_thread_zero_in_block()) {
quiet();
}
sync_all();
sync_all_wg();
__syncthreads();
}
@@ -134,7 +174,7 @@ __device__ void IPCContext::barrier(rocshmem_team_t team) {
if (is_thread_zero_in_block()) {
quiet();
}
internal_sync(pe, pe_start, pe_stride, pe_size, p_sync);
internal_sync_wg(pe, pe_start, pe_stride, pe_size, p_sync);
__syncthreads();
}
+3 -3
Просмотреть файл
@@ -468,7 +468,7 @@ __device__ void IPCContext::internal_broadcast(T *dst, const T *src, int nelems,
}
// Synchronize on completion of broadcast
internal_sync(my_pe, pe_start, stride, pe_size, p_sync);
internal_sync_wg(my_pe, pe_start, stride, pe_size, p_sync);
}
template <typename T>
@@ -497,7 +497,7 @@ __device__ void IPCContext::alltoall_linear(rocshmem_team_t team, T *dst,
quiet();
}
// wait until everyone has obtained their designated data
internal_sync(my_pe, pe_start, stride, pe_size, pSync);
internal_sync_wg(my_pe, pe_start, stride, pe_size, pSync);
}
template <typename T>
@@ -527,7 +527,7 @@ __device__ void IPCContext::fcollect_linear(rocshmem_team_t team, T *dst,
quiet();
}
// wait until everyone has obtained their designated data
internal_sync(my_pe, pe_start, stride, pe_size, pSync);
internal_sync_wg(my_pe, pe_start, stride, pe_size, pSync);
}
// Block/wave functions
+1 -1
Просмотреть файл
@@ -45,7 +45,7 @@ class IPCDefaultContextProxy {
size_t num_elems = 1)
: constructed_{true}, proxy_{num_elems} {
auto ctx{proxy_.get()};
new (ctx) IPCContext(reinterpret_cast<Backend*>(backend));
new (ctx) IPCContext(reinterpret_cast<Backend*>(backend), 0);
ctx->tinfo = tinfo;
rocshmem_ctx_t local{ctx, tinfo};
set_internal_ctx(&local);
+29 -1
Просмотреть файл
@@ -170,6 +170,20 @@ __device__ void *ROContext::shmem_ptr(const void *dest, int pe) {
}
__device__ void ROContext::barrier_all() {
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(), is_default_ctx);
}
__device__ void ROContext::barrier_all_wave() {
if (is_thread_zero_in_wave()) {
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(), is_default_ctx);
}
}
__device__ void ROContext::barrier_all_wg() {
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_BARRIER, nullptr, nullptr, 0, 0, 0, 0, 0,
nullptr, nullptr, (MPI_Comm)NULL, ro_net_win_id,
@@ -189,6 +203,20 @@ __device__ void ROContext::barrier(rocshmem_team_t team) {
}
__device__ void ROContext::sync_all() {
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(), is_default_ctx);
}
__device__ void ROContext::sync_all_wave() {
if (is_thread_zero_in_wave()) {
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(), is_default_ctx);
}
}
__device__ void ROContext::sync_all_wg() {
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0,
nullptr, nullptr, (MPI_Comm)NULL, ro_net_win_id,
@@ -197,7 +225,7 @@ __device__ void ROContext::sync_all() {
__syncthreads();
}
__device__ void ROContext::sync(rocshmem_team_t team) {
__device__ void ROContext::sync_wg(rocshmem_team_t team) {
ROTeam *team_obj = reinterpret_cast<ROTeam *>(team);
if (is_thread_zero_in_block()) {
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
+9 -1
Просмотреть файл
@@ -65,11 +65,19 @@ class ROContext : public Context {
__device__ void barrier_all();
__device__ void barrier_all_wave();
__device__ void barrier_all_wg();
__device__ void barrier(rocshmem_team_t team);
__device__ void sync_all();
__device__ void sync(rocshmem_team_t team);
__device__ void sync_all_wave();
__device__ void sync_all_wg();
__device__ void sync_wg(rocshmem_team_t team);
template <typename T>
__device__ void p(T *dest, T value, int pe);
+35 -3
Просмотреть файл
@@ -570,12 +570,24 @@ __device__ int rocshmem_test(T *ivars, int cmp, T val) {
return ctx_internal->test(ivars, cmp, val);
}
__device__ void rocshmem_ctx_wg_barrier_all(rocshmem_ctx_t ctx) {
__device__ void rocshmem_ctx_barrier_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_barrier_all\n");
get_internal_ctx(ctx)->barrier_all();
}
__device__ void rocshmem_ctx_wave_barrier_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_wave_barrier_all\n");
get_internal_ctx(ctx)->barrier_all_wave();
}
__device__ void rocshmem_ctx_wg_barrier_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_wg_barrier_all\n");
get_internal_ctx(ctx)->barrier_all_wg();
}
__device__ void rocshmem_wg_barrier_all() {
rocshmem_ctx_wg_barrier_all(ROCSHMEM_CTX_DEFAULT);
}
@@ -586,12 +598,32 @@ __device__ void rocshmem_barrier(rocshmem_team_t team) {
get_internal_ctx(ROCSHMEM_CTX_DEFAULT)->barrier(team);
}
__device__ void rocshmem_ctx_wg_sync_all(rocshmem_ctx_t ctx) {
__device__ void rocshmem_ctx_sync_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_sync_all\n");
get_internal_ctx(ctx)->sync_all();
}
__device__ void rocshmem_sync_all() {
rocshmem_ctx_sync_all(ROCSHMEM_CTX_DEFAULT);
}
__device__ void rocshmem_ctx_wave_sync_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_wave_sync_all\n");
get_internal_ctx(ctx)->sync_all_wave();
}
__device__ void rocshmem_wave_sync_all() {
rocshmem_ctx_wave_sync_all(ROCSHMEM_CTX_DEFAULT);
}
__device__ void rocshmem_ctx_wg_sync_all(rocshmem_ctx_t ctx) {
GPU_DPRINTF("Function: rocshmem_ctx_wg_sync_all\n");
get_internal_ctx(ctx)->sync_all_wg();
}
__device__ void rocshmem_wg_sync_all() {
rocshmem_ctx_wg_sync_all(ROCSHMEM_CTX_DEFAULT);
}
@@ -600,7 +632,7 @@ __device__ void rocshmem_ctx_wg_team_sync(rocshmem_ctx_t ctx,
rocshmem_team_t team) {
GPU_DPRINTF("Function: rocshmem_ctx_sync_all\n");
get_internal_ctx(ctx)->sync(team);
get_internal_ctx(ctx)->sync_wg(team);
}
__device__ void rocshmem_wg_team_sync(rocshmem_team_t team) {
+4
Просмотреть файл
@@ -43,6 +43,8 @@ enum rocshmem_stats {
NUM_QUIET,
NUM_TO_ALL,
NUM_BARRIER_ALL,
NUM_BARRIER_ALL_WAVE,
NUM_BARRIER_ALL_WG,
NUM_WAIT_UNTIL,
NUM_WAIT_UNTIL_ANY,
NUM_WAIT_UNTIL_ALL,
@@ -70,6 +72,8 @@ enum rocshmem_stats {
NUM_TEST,
NUM_SHMEM_PTR,
NUM_SYNC_ALL,
NUM_SYNC_ALL_WAVE,
NUM_SYNC_ALL_WG,
NUM_BROADCAST,
NUM_PUT_WG,
NUM_PUT_NBI_WG,