GDA get* APIs (#243)

* feat(GDA): add `get*` and `get*_nbi` APIs for mlx and bnxt NICs
   - implemented thread, wave and wg variants of `get*` and `get_nbi`.

* test(GDA): enable functional tests for `get*` and `get*_nbi` APIs

[ROCm/rocshmem commit: 671f8187f4]
This commit is contained in:
Avinash Kethineedi
2025-09-10 12:24:53 -04:00
committed by GitHub
parent 526784233b
commit 6860bc1275
5 changed files with 91 additions and 44 deletions
@@ -81,8 +81,20 @@ __device__ void GDAContext::putmem(void *dest, const void *source, size_t nelems
__device__ void GDAContext::getmem(void *dest, const void *source, size_t nelems,
int pe) {
printf("rocshmem::gda:getmem not implemented\n");
abort();
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
bool need_turn {true};
uint64_t turns = __ballot(need_turn);
while (turns) {
uint8_t lane = __ffsll((unsigned long long)turns) - 1;
int pe_turn = __shfl(pe, lane);
if (pe_turn == pe) {
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
qps[pe].quiet();
need_turn = false;
}
turns = __ballot(need_turn);
}
}
__device__ void GDAContext::putmem_nbi(void *dest, const void *source,
@@ -103,8 +115,19 @@ __device__ void GDAContext::putmem_nbi(void *dest, const void *source,
__device__ void GDAContext::getmem_nbi(void *dest, const void *source,
size_t nelems, int pe) {
printf("rocshmem::gda:getmem_nbi not implemented\n");
abort();
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
bool need_turn {true};
uint64_t turns = __ballot(need_turn);
while (turns) {
uint8_t lane = __ffsll((unsigned long long)turns) - 1;
int pe_turn = __shfl(pe, lane);
if (pe_turn == pe) {
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
need_turn = false;
}
turns = __ballot(need_turn);
}
}
__device__ void GDAContext::fence() { //TODO: optimize
@@ -139,9 +162,11 @@ __device__ void GDAContext::putmem_wg(void *dest, const void *source,
__device__ void GDAContext::getmem_wg(void *dest, const void *source,
size_t nelems, int pe) {
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
if (is_thread_zero_in_block()) {
printf("rocshmem::gda:getmem_wg not implemented\n");
abort();
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
qps[pe].quiet();
}
}
@@ -155,9 +180,10 @@ __device__ void GDAContext::putmem_nbi_wg(void *dest, const void *source,
__device__ void GDAContext::getmem_nbi_wg(void *dest, const void *source,
size_t nelems, int pe) {
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
if (is_thread_zero_in_block()) {
printf("rocshmem::gda:getmem_nbi_wg not implemented\n");
abort();
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
}
}
@@ -172,9 +198,11 @@ __device__ void GDAContext::putmem_wave(void *dest, const void *source,
__device__ void GDAContext::getmem_wave(void *dest, const void *source,
size_t nelems, int pe) {
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
if (is_thread_zero_in_wave()) {
printf("rocshmem::gda:getmem_wave not implemented\n");
abort();
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
qps[pe].quiet();
}
}
@@ -188,9 +216,10 @@ __device__ void GDAContext::putmem_nbi_wave(void *dest, const void *source,
__device__ void GDAContext::getmem_nbi_wave(void *dest, const void *source,
size_t nelems, int pe) {
const char *src_typed = reinterpret_cast<const char *>(source);
uint64_t L_offset = const_cast<char *>(src_typed) - base_heap[my_pe];
if (is_thread_zero_in_wave()) {
printf("rocshmem::gda:getmem_nbi_wave not implemented\n");
abort();
qps[pe].get_nbi(dest, base_heap[pe] + L_offset, nelems, pe);
}
}