Update collective APIs naming (#77)
* Update the naming convention for collective APIs to ensure consistency across the interface.
* Move all collective API declarations to rocshmem_COLL.hpp
* The following APIs were updated as part of this change:
- `barrier`
- `barrier_all`
- `sync`
- `sync_all`
- `all_to_all`
- `broadcast`
- `fcollect`
- `all_reduce`
* Update header file generation code for collective APIs
[ROCm/rocshmem commit: 68421895d6]
This commit is contained in:
committato da
GitHub
parent
5b22ddd1ff
commit
41d5d739e2
@@ -53,11 +53,11 @@ __global__ void BarrierAllTest(int loop, int skip, long long int *start_time,
|
||||
break;
|
||||
case WAVEBarrierAllTestType:
|
||||
if(wf_id == 0) {
|
||||
rocshmem_ctx_wave_barrier_all(ctx);
|
||||
rocshmem_ctx_barrier_all_wave(ctx);
|
||||
}
|
||||
break;
|
||||
case WGBarrierAllTestType:
|
||||
rocshmem_ctx_wg_barrier_all(ctx);
|
||||
rocshmem_ctx_barrier_all_wg(ctx);
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
|
||||
@@ -53,11 +53,11 @@ __global__ void SyncAllTest(int loop, int skip, long long int *start_time,
|
||||
break;
|
||||
case WAVESyncAllTestType:
|
||||
if(wf_id == 0) {
|
||||
rocshmem_ctx_wave_sync_all(ctx);
|
||||
rocshmem_ctx_sync_all_wave(ctx);
|
||||
}
|
||||
break;
|
||||
case WGSyncAllTestType:
|
||||
rocshmem_ctx_wg_sync_all(ctx);
|
||||
rocshmem_ctx_sync_all_wg(ctx);
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
|
||||
@@ -45,16 +45,16 @@ __global__ void SyncTest(int loop, int skip, long long int *start_time,
|
||||
switch (type) {
|
||||
case SyncTestType:
|
||||
if(t_id == 0) {
|
||||
rocshmem_ctx_team_sync(ctx, teams[wg_id]);
|
||||
rocshmem_ctx_sync(ctx, teams[wg_id]);
|
||||
}
|
||||
break;
|
||||
case WAVESyncTestType:
|
||||
if(wf_id == 0) {
|
||||
rocshmem_ctx_wave_team_sync(ctx, teams[wg_id]);
|
||||
rocshmem_ctx_sync_wave(ctx, teams[wg_id]);
|
||||
}
|
||||
break;
|
||||
case WGSyncTestType:
|
||||
rocshmem_ctx_wg_team_sync(ctx, teams[wg_id]);
|
||||
rocshmem_ctx_sync_wg(ctx, teams[wg_id]);
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
|
||||
@@ -32,7 +32,7 @@ __device__ void wg_team_alltoall(rocshmem_ctx_t ctx, rocshmem_team_t team,
|
||||
template <> \
|
||||
__device__ void wg_team_alltoall<T>(rocshmem_ctx_t ctx, rocshmem_team_t team,\
|
||||
T * dest, const T *source, int nelem) { \
|
||||
rocshmem_ctx_##TNAME##_wg_alltoall(ctx, team, dest, source, nelem); \
|
||||
rocshmem_ctx_##TNAME##_alltoall_wg(ctx, team, dest, source, nelem); \
|
||||
}
|
||||
|
||||
TEAM_ALLTOALL_DEF_GEN(float, float)
|
||||
|
||||
@@ -51,11 +51,11 @@ __global__ void TeamBarrierTest(int loop, int skip, long long int *start_time,
|
||||
break;
|
||||
case TeamWAVEBarrierTestType:
|
||||
if(wf_id == 0) {
|
||||
rocshmem_ctx_wave_barrier(ctx, teams[wg_id]);
|
||||
rocshmem_ctx_barrier_wave(ctx, teams[wg_id]);
|
||||
}
|
||||
break;
|
||||
case TeamWGBarrierTestType:
|
||||
rocshmem_ctx_wg_barrier(ctx, teams[wg_id]);
|
||||
rocshmem_ctx_barrier_wg(ctx, teams[wg_id]);
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
|
||||
@@ -34,7 +34,7 @@ __device__ void wg_team_broadcast(rocshmem_ctx_t ctx, rocshmem_team_t team,
|
||||
__device__ void wg_team_broadcast<T>( \
|
||||
rocshmem_ctx_t ctx, rocshmem_team_t team, T * dest, const T *source, \
|
||||
int nelem, int pe_root) { \
|
||||
rocshmem_ctx_##TNAME##_wg_broadcast(ctx, team, dest, source, nelem, \
|
||||
rocshmem_ctx_##TNAME##_broadcast_wg(ctx, team, dest, source, nelem, \
|
||||
pe_root); \
|
||||
}
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ __device__ void wg_team_fcollect(rocshmem_ctx_t ctx, rocshmem_team_t team,
|
||||
template <> \
|
||||
__device__ void wg_team_fcollect<T>(rocshmem_ctx_t ctx, rocshmem_team_t team,\
|
||||
T * dest, const T *source, int nelem) { \
|
||||
rocshmem_ctx_##TNAME##_wg_fcollect(ctx, team, dest, source, nelem); \
|
||||
rocshmem_ctx_##TNAME##_fcollect_wg(ctx, team, dest, source, nelem); \
|
||||
}
|
||||
|
||||
TEAM_FCOLLECT_DEF_GEN(float, float)
|
||||
|
||||
@@ -35,7 +35,7 @@ __device__ int wg_team_reduce(rocshmem_ctx_t ctx, rocshmem_team_t, T *dest,
|
||||
__device__ int wg_team_reduce<T, Op>(rocshmem_ctx_t ctx, \
|
||||
rocshmem_team_t team, T * dest, \
|
||||
const T *source, int nreduce) { \
|
||||
return rocshmem_ctx_##TNAME##_##Op_API##_wg_reduce(ctx, team, dest, \
|
||||
return rocshmem_ctx_##TNAME##_##Op_API##_reduce_wg(ctx, team, dest, \
|
||||
source, nreduce); \
|
||||
}
|
||||
|
||||
@@ -93,7 +93,7 @@ __global__ void TeamReductionTest(int loop, int skip, long long int *start_time,
|
||||
start_time[wg_id] = wall_clock64();
|
||||
}
|
||||
wg_team_reduce<T1, T2>(ctx, team, r_buf, s_buf, size);
|
||||
rocshmem_ctx_wg_barrier_all(ctx);
|
||||
rocshmem_ctx_barrier_all_wg(ctx);
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
Fai riferimento in un nuovo problema
Block a user