Use new naming scheme
Этот коммит содержится в:
@@ -24,20 +24,20 @@ using namespace rocshmem;
|
||||
|
||||
/* Declare the template with a generic implementation */
|
||||
template <typename T>
|
||||
__device__ void wg_team_broadcast(roc_shmem_ctx_t ctx, roc_shmem_team_t team,
|
||||
__device__ void wg_team_broadcast(rocshmem_ctx_t ctx, rocshmem_team_t team,
|
||||
T *dest, const T *source, int nelem,
|
||||
int pe_root) {
|
||||
return;
|
||||
}
|
||||
|
||||
/* Define templates to call ROC_SHMEM */
|
||||
#define TEAM_BROADCAST_DEF_GEN(T, TNAME) \
|
||||
template <> \
|
||||
__device__ void wg_team_broadcast<T>( \
|
||||
roc_shmem_ctx_t ctx, roc_shmem_team_t team, T * dest, const T *source, \
|
||||
int nelem, int pe_root) { \
|
||||
roc_shmem_ctx_##TNAME##_wg_broadcast(ctx, team, dest, source, nelem, \
|
||||
pe_root); \
|
||||
/* Define templates to call ROCSHMEM */
|
||||
#define TEAM_BROADCAST_DEF_GEN(T, TNAME) \
|
||||
template <> \
|
||||
__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, \
|
||||
pe_root); \
|
||||
}
|
||||
|
||||
TEAM_BROADCAST_DEF_GEN(float, float)
|
||||
@@ -55,7 +55,7 @@ TEAM_BROADCAST_DEF_GEN(unsigned int, uint)
|
||||
TEAM_BROADCAST_DEF_GEN(unsigned long, ulong)
|
||||
TEAM_BROADCAST_DEF_GEN(unsigned long long, ulonglong)
|
||||
|
||||
roc_shmem_team_t team_bcast_world_dup;
|
||||
rocshmem_team_t team_bcast_world_dup;
|
||||
|
||||
/******************************************************************************
|
||||
* DEVICE TEST KERNEL
|
||||
@@ -64,20 +64,20 @@ template <typename T1>
|
||||
__global__ void TeamBroadcastTest(int loop, int skip, uint64_t *timer,
|
||||
T1 *source_buf, T1 *dest_buf, int size,
|
||||
ShmemContextType ctx_type,
|
||||
roc_shmem_team_t team) {
|
||||
__shared__ roc_shmem_ctx_t ctx;
|
||||
rocshmem_team_t team) {
|
||||
__shared__ rocshmem_ctx_t ctx;
|
||||
|
||||
roc_shmem_wg_init();
|
||||
roc_shmem_wg_ctx_create(ctx_type, &ctx);
|
||||
rocshmem_wg_init();
|
||||
rocshmem_wg_ctx_create(ctx_type, &ctx);
|
||||
|
||||
int n_pes = roc_shmem_ctx_n_pes(ctx);
|
||||
int n_pes = rocshmem_ctx_n_pes(ctx);
|
||||
|
||||
__syncthreads();
|
||||
|
||||
uint64_t start;
|
||||
for (int i = 0; i < loop; i++) {
|
||||
if (i == skip && hipThreadIdx_x == 0) {
|
||||
start = roc_shmem_timer();
|
||||
start = rocshmem_timer();
|
||||
}
|
||||
|
||||
wg_team_broadcast<T1>(ctx, team,
|
||||
@@ -85,17 +85,17 @@ __global__ void TeamBroadcastTest(int loop, int skip, uint64_t *timer,
|
||||
source_buf, // const T* source
|
||||
size, // int nelement
|
||||
0); // int PE_root
|
||||
roc_shmem_ctx_wg_barrier_all(ctx);
|
||||
rocshmem_ctx_wg_barrier_all(ctx);
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
if (hipThreadIdx_x == 0) {
|
||||
timer[hipBlockIdx_x] = roc_shmem_timer() - start;
|
||||
timer[hipBlockIdx_x] = rocshmem_timer() - start;
|
||||
}
|
||||
|
||||
roc_shmem_wg_ctx_destroy(&ctx);
|
||||
roc_shmem_wg_finalize();
|
||||
rocshmem_wg_ctx_destroy(&ctx);
|
||||
rocshmem_wg_finalize();
|
||||
}
|
||||
|
||||
/******************************************************************************
|
||||
@@ -106,22 +106,22 @@ TeamBroadcastTester<T1>::TeamBroadcastTester(
|
||||
TesterArguments args, std::function<void(T1 &, T1 &)> f1,
|
||||
std::function<std::pair<bool, std::string>(const T1 &)> f2)
|
||||
: Tester(args), init_buf{f1}, verify_buf{f2} {
|
||||
source_buf = (T1 *)roc_shmem_malloc(args.max_msg_size * sizeof(T1));
|
||||
dest_buf = (T1 *)roc_shmem_malloc(args.max_msg_size * sizeof(T1));
|
||||
source_buf = (T1 *)rocshmem_malloc(args.max_msg_size * sizeof(T1));
|
||||
dest_buf = (T1 *)rocshmem_malloc(args.max_msg_size * sizeof(T1));
|
||||
}
|
||||
|
||||
template <typename T1>
|
||||
TeamBroadcastTester<T1>::~TeamBroadcastTester() {
|
||||
roc_shmem_free(source_buf);
|
||||
roc_shmem_free(dest_buf);
|
||||
rocshmem_free(source_buf);
|
||||
rocshmem_free(dest_buf);
|
||||
}
|
||||
|
||||
template <typename T1>
|
||||
void TeamBroadcastTester<T1>::preLaunchKernel() {
|
||||
int n_pes = roc_shmem_team_n_pes(ROC_SHMEM_TEAM_WORLD);
|
||||
int n_pes = rocshmem_team_n_pes(ROCSHMEM_TEAM_WORLD);
|
||||
|
||||
team_bcast_world_dup = ROC_SHMEM_TEAM_INVALID;
|
||||
roc_shmem_team_split_strided(ROC_SHMEM_TEAM_WORLD, 0, 1, n_pes, nullptr, 0,
|
||||
team_bcast_world_dup = ROCSHMEM_TEAM_INVALID;
|
||||
rocshmem_team_split_strided(ROCSHMEM_TEAM_WORLD, 0, 1, n_pes, nullptr, 0,
|
||||
&team_bcast_world_dup);
|
||||
}
|
||||
|
||||
@@ -140,7 +140,7 @@ void TeamBroadcastTester<T1>::launchKernel(dim3 gridSize, dim3 blockSize,
|
||||
|
||||
template <typename T1>
|
||||
void TeamBroadcastTester<T1>::postLaunchKernel() {
|
||||
roc_shmem_team_destroy(team_bcast_world_dup);
|
||||
rocshmem_team_destroy(team_bcast_world_dup);
|
||||
}
|
||||
|
||||
template <typename T1>
|
||||
|
||||
Ссылка в новой задаче
Block a user