diff --git a/projects/rocshmem/src/ipc/context_ipc_device.cpp b/projects/rocshmem/src/ipc/context_ipc_device.cpp index ede7633b06..6fbb362d73 100644 --- a/projects/rocshmem/src/ipc/context_ipc_device.cpp +++ b/projects/rocshmem/src/ipc/context_ipc_device.cpp @@ -39,7 +39,8 @@ namespace rocshmem { __host__ IPCContext::IPCContext(Backend *b) : Context(b, false) { IPCBackend *backend{static_cast(b)}; - ipcImpl = &backend->ipcImpl; + ipcImpl_.ipc_bases = b->ipcImpl.ipc_bases; + ipcImpl_.shm_size = b->ipcImpl.shm_size; auto *bp{backend->ipc_backend_proxy.get()}; @@ -59,22 +60,18 @@ __device__ void IPCContext::ctx_destroy(){ __device__ void IPCContext::putmem(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; uint64_t L_offset = - reinterpret_cast(dest) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy(ipcImpl->ipc_bases[local_pe] + L_offset, + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy(ipcImpl_.ipc_bases[pe] + L_offset, const_cast(source), nelems); } __device__ void IPCContext::getmem(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; const char *src_typed = reinterpret_cast(source); uint64_t L_offset = - const_cast(src_typed) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy(dest, ipcImpl->ipc_bases[local_pe] + L_offset, nelems); + const_cast(src_typed) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy(dest, ipcImpl_.ipc_bases[pe] + L_offset, nelems); } __device__ void IPCContext::putmem_nbi(void *dest, const void *source, @@ -103,23 +100,19 @@ __device__ void *IPCContext::shmem_ptr(const void *dest, int pe) { __device__ void IPCContext::putmem_wg(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; uint64_t L_offset = - reinterpret_cast(dest) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy_wg(ipcImpl->ipc_bases[local_pe] + L_offset, + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy_wg(ipcImpl_.ipc_bases[pe] + L_offset, const_cast(source), nelems); __syncthreads(); } __device__ void IPCContext::getmem_wg(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; const char *src_typed = reinterpret_cast(source); uint64_t L_offset = - const_cast(src_typed) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy_wg(dest, ipcImpl->ipc_bases[local_pe] + L_offset, nelems); + const_cast(src_typed) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy_wg(dest, ipcImpl_.ipc_bases[pe] + L_offset, nelems); __syncthreads(); } @@ -135,22 +128,18 @@ __device__ void IPCContext::getmem_nbi_wg(void *dest, const void *source, __device__ void IPCContext::putmem_wave(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; uint64_t L_offset = - reinterpret_cast(dest) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy_wave(ipcImpl->ipc_bases[local_pe] + L_offset, + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy_wave(ipcImpl_.ipc_bases[pe] + L_offset, const_cast(source), nelems); } __device__ void IPCContext::getmem_wave(void *dest, const void *source, size_t nelems, int pe) { - // TODO (Avinash) check if PE is available for IPC using (isIpcAvailable) - int local_pe = pe % ipcImpl->shm_size; const char *src_typed = reinterpret_cast(source); uint64_t L_offset = - const_cast(src_typed) - ipcImpl->ipc_bases[my_pe]; - ipcImpl->ipcCopy_wave(dest, ipcImpl->ipc_bases[local_pe] + L_offset, + const_cast(src_typed) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcCopy_wave(dest, ipcImpl_.ipc_bases[pe] + L_offset, nelems); } diff --git a/projects/rocshmem/src/ipc/context_ipc_tmpl_device.hpp b/projects/rocshmem/src/ipc/context_ipc_tmpl_device.hpp index 81ec0b2b97..91bdbd45e7 100644 --- a/projects/rocshmem/src/ipc/context_ipc_tmpl_device.hpp +++ b/projects/rocshmem/src/ipc/context_ipc_tmpl_device.hpp @@ -71,13 +71,19 @@ __device__ void IPCContext::get_nbi(T *dest, const T *source, size_t nelems, // Atomics template -__device__ void IPCContext::amo_add(void *dst, T value, int pe) { - assert(false); +__device__ void IPCContext::amo_add(void *dest, T value, int pe) { + uint64_t L_offset = + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcAMOAdd( + reinterpret_cast(ipcImpl_.ipc_bases[pe] + L_offset), value); } template -__device__ void IPCContext::amo_set(void *dst, T value, int pe) { - assert(false); +__device__ void IPCContext::amo_set(void *dest, T value, int pe) { + uint64_t L_offset = + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcAMOSet( + reinterpret_cast(ipcImpl_.ipc_bases[pe] + L_offset), value); } template @@ -120,20 +126,29 @@ __device__ void IPCContext::amo_xor(void *dst, T value, int pe) { } template -__device__ void IPCContext::amo_cas(void *dst, T value, T cond, int pe) { - assert(false); +__device__ void IPCContext::amo_cas(void *dest, T value, T cond, int pe) { + uint64_t L_offset = + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + ipcImpl_.ipcAMOCas( + reinterpret_cast(ipcImpl_.ipc_bases[pe] + L_offset), cond, + value); } template -__device__ T IPCContext::amo_fetch_add(void *dst, T value, int pe) { - assert(false); - return 0; +__device__ T IPCContext::amo_fetch_add(void *dest, T value, int pe) { + uint64_t L_offset = + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + return ipcImpl_.ipcAMOFetchAdd( + reinterpret_cast(ipcImpl_.ipc_bases[pe] + L_offset), value); } template -__device__ T IPCContext::amo_fetch_cas(void *dst, T value, T cond, int pe) { - assert(false); - return 0; +__device__ T IPCContext::amo_fetch_cas(void *dest, T value, T cond, int pe) { + uint64_t L_offset = + reinterpret_cast(dest) - ipcImpl_.ipc_bases[my_pe]; + return ipcImpl_.ipcAMOFetchCas( + reinterpret_cast(ipcImpl_.ipc_bases[pe] + L_offset), cond, + value); } // Collectives