Implement default RO context (#64)
* Allocate default context buffers and initialize queue for management - Allocated the status flag, g return, and atomic return buffers for the default context. - Initialized `AtomicWFQueueProxy` instances to manage these buffers efficiently for concurrent access. * Update `BlockHandle` with default context buffers * Add default context flag and update buffer retrieval functions - Added a flag to distinguish the default context from other contexts. - Modified return buffer functionns and `get_status_flag` function to accommodate the default context * Add default context primitive tests - get, put, get_nbi, put_nbi, g, and p APIs.
This commit is contained in:
committed by
GitHub
parent
4e48c9748e
commit
867519e1d0
@@ -41,10 +41,14 @@
|
||||
|
||||
namespace rocshmem {
|
||||
|
||||
__host__ ROContext::ROContext(Backend *b, size_t block_id)
|
||||
: Context(b, false) {
|
||||
__host__ ROContext::ROContext(Backend *b,
|
||||
size_t block_id,
|
||||
bool default_ctx)
|
||||
: Context(b, false),
|
||||
is_default_ctx{default_ctx} {
|
||||
|
||||
ROBackend *backend{static_cast<ROBackend *>(b)};
|
||||
if (block_id == -1) {
|
||||
if (is_default_ctx) {
|
||||
block_handle = backend->default_block_handle_proxy_.get();
|
||||
} else {
|
||||
auto block_base{backend->block_handle_proxy_.get()};
|
||||
@@ -72,7 +76,8 @@ __device__ void ROContext::putmem(void *dest, const void *source, size_t nelems,
|
||||
}
|
||||
build_queue_element(RO_NET_PUT, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,7 +96,8 @@ __device__ void ROContext::getmem(void *dest, const void *source, size_t nelems,
|
||||
}
|
||||
build_queue_element(RO_NET_GET, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -136,20 +142,20 @@ __device__ void ROContext::getmem_nbi(void *dest, const void *source,
|
||||
__device__ void ROContext::fence() {
|
||||
build_queue_element(RO_NET_FENCE, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
|
||||
nullptr, (MPI_Comm)NULL, ro_net_win_id, block_handle,
|
||||
true, get_status_flag());
|
||||
true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
|
||||
__device__ void ROContext::fence(int pe) {
|
||||
// TODO(khamidou): need to check if per pe has any special handling
|
||||
build_queue_element(RO_NET_FENCE, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
|
||||
nullptr, (MPI_Comm)NULL, ro_net_win_id, block_handle,
|
||||
true, get_status_flag());
|
||||
true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
|
||||
__device__ void ROContext::quiet() {
|
||||
build_queue_element(RO_NET_QUIET, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
|
||||
nullptr, (MPI_Comm)NULL, ro_net_win_id, block_handle,
|
||||
true, get_status_flag());
|
||||
true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
|
||||
__device__ void *ROContext::shmem_ptr(const void *dest, int pe) {
|
||||
@@ -167,7 +173,7 @@ __device__ void ROContext::barrier_all() {
|
||||
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,
|
||||
block_handle, true, get_status_flag());
|
||||
block_handle, true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
@@ -177,7 +183,7 @@ __device__ void ROContext::barrier(rocshmem_team_t team) {
|
||||
if (is_thread_zero_in_block()) {
|
||||
build_queue_element(RO_NET_BARRIER, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
|
||||
nullptr, team_obj->mpi_comm, ro_net_win_id, block_handle,
|
||||
true, get_status_flag());
|
||||
true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
@@ -186,7 +192,7 @@ __device__ void ROContext::sync_all() {
|
||||
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,
|
||||
block_handle, true, get_status_flag());
|
||||
block_handle, true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
@@ -196,7 +202,7 @@ __device__ void ROContext::sync(rocshmem_team_t team) {
|
||||
if (is_thread_zero_in_block()) {
|
||||
build_queue_element(RO_NET_SYNC, nullptr, nullptr, 0, 0, 0, 0, 0, nullptr,
|
||||
nullptr, team_obj->mpi_comm, ro_net_win_id, block_handle,
|
||||
true, get_status_flag());
|
||||
true, get_status_flag(), is_default_ctx);
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
@@ -209,7 +215,7 @@ __device__ void ROContext::ctx_destroy() {
|
||||
|
||||
build_queue_element(RO_NET_FINALIZE, nullptr, nullptr, 0, 0, 0, 0, 0,
|
||||
nullptr, nullptr, (MPI_Comm)NULL, ro_net_win_id,
|
||||
block_handle, true, get_status_flag());
|
||||
block_handle, true, get_status_flag(), is_default_ctx);
|
||||
|
||||
int buffer_id = ro_net_win_id;
|
||||
backend->queue_.descriptor(buffer_id)->write_index = block_handle->write_index;
|
||||
@@ -233,7 +239,8 @@ __device__ void ROContext::putmem_wg(void *dest, const void *source,
|
||||
if (is_thread_zero_in_block()) {
|
||||
build_queue_element(RO_NET_PUT, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
@@ -251,7 +258,8 @@ __device__ void ROContext::getmem_wg(void *dest, const void *source,
|
||||
if (is_thread_zero_in_block()) {
|
||||
build_queue_element(RO_NET_GET, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
@@ -305,7 +313,8 @@ __device__ void ROContext::putmem_wave(void *dest, const void *source,
|
||||
if (is_thread_zero_in_wave()) {
|
||||
build_queue_element(RO_NET_PUT, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -323,7 +332,8 @@ __device__ void ROContext::getmem_wave(void *dest, const void *source,
|
||||
if (is_thread_zero_in_wave()) {
|
||||
build_queue_element(RO_NET_GET, dest, const_cast<void *>(source), nelems,
|
||||
pe, 0, 0, 0, nullptr, nullptr, (MPI_Comm)NULL,
|
||||
ro_net_win_id, block_handle, true, get_status_flag());
|
||||
ro_net_win_id, block_handle, true, get_status_flag(),
|
||||
is_default_ctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -592,7 +602,7 @@ __device__ void build_queue_element(
|
||||
ro_net_cmds type, void *dst, void *src, size_t size, int pe,
|
||||
int logPE_stride, int PE_size, int PE_root, void *pWrk, long *pSync,
|
||||
MPI_Comm team_comm, int ro_net_win_id, BlockHandle *handle,
|
||||
bool blocking, volatile char *status, ROCSHMEM_OP op,
|
||||
bool blocking, volatile char *status, bool default_ctx, ROCSHMEM_OP op,
|
||||
ro_net_types datatype) {
|
||||
auto write_slot{next_write_slot(handle)};
|
||||
auto queue_element = &handle->queue[write_slot];
|
||||
@@ -661,25 +671,49 @@ __device__ void build_queue_element(
|
||||
*(queue_element->status) = 0;
|
||||
__threadfence();
|
||||
}
|
||||
if (default_ctx) {
|
||||
handle->default_ctx_status->enqueue(status);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ uint64_t *ROContext::get_atomic_ret_buf() {
|
||||
uint64_t *atomic_base_ptr{
|
||||
reinterpret_cast<uint64_t*>(block_handle->atomic_ret)};
|
||||
int thread_id{get_flat_block_id()};
|
||||
return &atomic_base_ptr[thread_id];
|
||||
uint64_t *buf_addr{nullptr};
|
||||
if(is_default_ctx) {
|
||||
buf_addr = block_handle->default_ctx_atomic_ret->dequeue();
|
||||
}
|
||||
else {
|
||||
uint64_t *atomic_base_ptr{
|
||||
reinterpret_cast<uint64_t*>(block_handle->atomic_ret)};
|
||||
int thread_id{get_flat_block_id()};
|
||||
buf_addr = &atomic_base_ptr[thread_id];
|
||||
}
|
||||
return buf_addr;
|
||||
}
|
||||
|
||||
__device__ uint64_t *ROContext::get_g_ret_buf() {
|
||||
uint64_t *g_ret{reinterpret_cast<uint64_t*>(block_handle->g_ret)};
|
||||
int thread_id{get_flat_block_id()};
|
||||
return &g_ret[thread_id];
|
||||
uint64_t *buf_addr{nullptr};
|
||||
if(is_default_ctx) {
|
||||
buf_addr = block_handle->default_ctx_g_ret->dequeue();
|
||||
}
|
||||
else {
|
||||
uint64_t *g_ret{reinterpret_cast<uint64_t*>(block_handle->g_ret)};
|
||||
int thread_id{get_flat_block_id()};
|
||||
buf_addr = &g_ret[thread_id];
|
||||
}
|
||||
return buf_addr;
|
||||
}
|
||||
|
||||
__device__ volatile char *ROContext::get_status_flag() {
|
||||
volatile char* status{block_handle->status};
|
||||
int thread_id{get_flat_block_id()};
|
||||
return &status[thread_id];
|
||||
volatile char *status_addr{nullptr};
|
||||
if(is_default_ctx) {
|
||||
status_addr = block_handle->default_ctx_status->dequeue();
|
||||
}
|
||||
else {
|
||||
volatile char* status{block_handle->status};
|
||||
int thread_id{get_flat_block_id()};
|
||||
status_addr = &status[thread_id];
|
||||
}
|
||||
return status_addr;
|
||||
}
|
||||
|
||||
} // namespace rocshmem
|
||||
|
||||
Reference in New Issue
Block a user