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:
Avinash Kethineedi
2025-03-25 18:51:54 -05:00
committed by GitHub
parent 4e48c9748e
commit 867519e1d0
12 changed files with 465 additions and 46 deletions
+62 -28
View File
@@ -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