@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user