Use new naming scheme

[ROCm/rocshmem commit: fd8dbc7fb6]
This commit is contained in:
Brandon Potter
2024-11-25 14:12:15 -06:00
parent 07088f0f2b
commit 913ce47ef1
179 changed files with 5250 additions and 5251 deletions
+3 -3
View File
@@ -83,13 +83,13 @@ int HostInterface::find_win_info_in_pool(WindowInfo* window_info) {
}
__host__ HostInterface::HostInterface(HdpPolicy* hdp_policy,
MPI_Comm roc_shmem_comm,
MPI_Comm rocshmem_comm,
SymmetricHeap* heap) {
/*
* Duplicate a communicator from roc_shem's comm
* world for the host interface
*/
MPI_Comm_dup(roc_shmem_comm, &host_comm_world_);
MPI_Comm_dup(rocshmem_comm, &host_comm_world_);
MPI_Comm_rank(host_comm_world_, &my_pe_);
MPI_Comm_rank(host_comm_world_, &num_pes_);
@@ -103,7 +103,7 @@ __host__ HostInterface::HostInterface(HdpPolicy* hdp_policy,
* Allocate and initialize pool of windows for contexts
*/
char* value{nullptr};
if ((value = getenv("ROC_SHMEM_MAX_NUM_HOST_CONTEXTS"))) {
if ((value = getenv("ROCSHMEM_MAX_NUM_HOST_CONTEXTS"))) {
max_num_ctxs_ = atoi(value);
}
+8 -8
View File
@@ -36,7 +36,7 @@
#include <map>
#include "roc_shmem/roc_shmem.hpp"
#include "rocshmem/rocshmem.hpp"
#include "../hdp_policy.hpp"
#include "../memory/symmetric_heap.hpp"
#include "../memory/window_info.hpp"
@@ -104,7 +104,7 @@ class HostInterface {
/**
* @brief Primary constructor
*/
__host__ HostInterface(HdpPolicy* hdp_policy, MPI_Comm roc_shmem_comm,
__host__ HostInterface(HdpPolicy* hdp_policy, MPI_Comm rocshmem_comm,
SymmetricHeap* heap);
/**
@@ -198,16 +198,16 @@ class HostInterface {
long* p_sync); // NOLINT(runtime/int)
template <typename T>
__host__ void broadcast(roc_shmem_team_t team, T* dest, const T* source,
__host__ void broadcast(rocshmem_team_t team, T* dest, const T* source,
int nelems, int pe_root);
template <typename T, ROC_SHMEM_OP Op>
template <typename T, ROCSHMEM_OP Op>
__host__ void to_all(T* dest, const T* source, int nreduce, int pe_start,
int log_pe_stride, int pe_size, T* p_wrk,
long* p_sync); // NOLINT(runtime/int)
template <typename T, ROC_SHMEM_OP Op>
__host__ int reduce(roc_shmem_team_t team, T* dest, const T* source, int nreduce);
template <typename T, ROCSHMEM_OP Op>
__host__ int reduce(rocshmem_team_t team, T* dest, const T* source, int nreduce);
template <typename T>
__host__ void wait_until(T *ivars, int cmp, T val,
@@ -288,7 +288,7 @@ class HostInterface {
__host__ MPI_Comm get_mpi_comm(int pe_start, int log_pe_stride, int pe_size);
__host__ MPI_Op get_mpi_op(ROC_SHMEM_OP Op);
__host__ MPI_Op get_mpi_op(ROCSHMEM_OP Op);
template <typename T>
__host__ MPI_Datatype get_mpi_type();
@@ -300,7 +300,7 @@ class HostInterface {
__host__ int test_and_compare(MPI_Aint offset, MPI_Datatype mpi_type,
int cmp, T val, MPI_Win win);
template <typename T, ROC_SHMEM_OP Op>
template <typename T, ROCSHMEM_OP Op>
__host__ void to_all_internal(MPI_Comm mpi_comm, T* dest, const T* source,
int nreduce);
+22 -22
View File
@@ -200,7 +200,7 @@ __host__ void HostInterface::broadcast(T* dest, const T* source, int nelems,
}
template <typename T>
__host__ void HostInterface::broadcast(roc_shmem_team_t team, T* dest,
__host__ void HostInterface::broadcast(rocshmem_team_t team, T* dest,
const T* source, int nelems,
int pe_root) {
DPRINTF("Function: Team-based host_broadcast\n");
@@ -216,24 +216,24 @@ __host__ void HostInterface::broadcast(roc_shmem_team_t team, T* dest,
return;
}
__host__ inline MPI_Op HostInterface::get_mpi_op(ROC_SHMEM_OP Op) {
__host__ inline MPI_Op HostInterface::get_mpi_op(ROCSHMEM_OP Op) {
switch (Op) {
case ROC_SHMEM_SUM:
case ROCSHMEM_SUM:
return MPI_SUM;
case ROC_SHMEM_MAX:
case ROCSHMEM_MAX:
return MPI_MAX;
case ROC_SHMEM_MIN:
case ROCSHMEM_MIN:
return MPI_MIN;
case ROC_SHMEM_PROD:
case ROCSHMEM_PROD:
return MPI_PROD;
case ROC_SHMEM_AND:
case ROCSHMEM_AND:
return MPI_BAND;
case ROC_SHMEM_OR:
case ROCSHMEM_OR:
return MPI_BOR;
case ROC_SHMEM_XOR:
case ROCSHMEM_XOR:
return MPI_BXOR;
default:
fprintf(stderr, "Unknown ROC_SHMEM op MPI conversion %d\n", Op);
fprintf(stderr, "Unknown rocSHMEM op MPI conversion %d\n", Op);
abort();
return 0;
}
@@ -330,7 +330,7 @@ __host__ T HostInterface::amo_fetch_cas(void* dst, T value, T cond, int pe,
return ret;
}
template <typename T, ROC_SHMEM_OP Op>
template <typename T, ROCSHMEM_OP Op>
__host__ void HostInterface::to_all_internal(MPI_Comm mpi_comm, T* dest,
const T* source, int nreduce) {
DPRINTF("Function: host_to_all_internal\n");
@@ -356,7 +356,7 @@ __host__ void HostInterface::to_all_internal(MPI_Comm mpi_comm, T* dest,
return;
}
template <typename T, ROC_SHMEM_OP Op>
template <typename T, ROCSHMEM_OP Op>
__host__ void HostInterface::to_all(T* dest, const T* source, int nreduce,
int pe_start, int log_pe_stride,
int pe_size, [[maybe_unused]] T* p_wrk,
@@ -375,8 +375,8 @@ __host__ void HostInterface::to_all(T* dest, const T* source, int nreduce,
return;
}
template <typename T, ROC_SHMEM_OP Op>
__host__ int HostInterface::reduce(roc_shmem_team_t team, T* dest,
template <typename T, ROCSHMEM_OP Op>
__host__ int HostInterface::reduce(rocshmem_team_t team, T* dest,
const T* source, int nreduce) {
DPRINTF("Function: Team-based host_reduce\n");
@@ -388,7 +388,7 @@ __host__ int HostInterface::reduce(roc_shmem_team_t team, T* dest,
to_all_internal<T, Op>(mpi_comm, dest, source, nreduce);
return ROC_SHMEM_SUCCESS;
return ROCSHMEM_SUCCESS;
}
template <typename T>
@@ -397,26 +397,26 @@ __host__ inline int HostInterface::compare(int cmp, T input_val,
int cond_satisfied{0};
switch (cmp) {
case ROC_SHMEM_CMP_EQ:
case ROCSHMEM_CMP_EQ:
cond_satisfied = (input_val == target_val) ? 1 : 0;
break;
case ROC_SHMEM_CMP_NE:
case ROCSHMEM_CMP_NE:
cond_satisfied = (input_val != target_val) ? 1 : 0;
break;
case ROC_SHMEM_CMP_GT:
case ROCSHMEM_CMP_GT:
cond_satisfied = (input_val > target_val) ? 1 : 0;
break;
case ROC_SHMEM_CMP_GE:
case ROCSHMEM_CMP_GE:
cond_satisfied = (input_val >= target_val) ? 1 : 0;
break;
case ROC_SHMEM_CMP_LT:
case ROCSHMEM_CMP_LT:
cond_satisfied = (input_val < target_val) ? 1 : 0;
break;
case ROC_SHMEM_CMP_LE:
case ROCSHMEM_CMP_LE:
cond_satisfied = (input_val <= target_val) ? 1 : 0;
break;
default:
assert(cmp >= ROC_SHMEM_CMP_EQ && cmp <= ROC_SHMEM_CMP_LE);
assert(cmp >= ROCSHMEM_CMP_EQ && cmp <= ROCSHMEM_CMP_LE);
break;
}