[GDA/BNXT] Optimize Alltoall using put signal (#334)

* Modularize bnxt

* add post_wqe_amo_single

* add alltoall with putsignal impl

* make ringing the doorbell optional

[ROCm/rocshmem commit: baaf8091b5]
This commit is contained in:
Yiltan
2025-12-05 12:41:22 -05:00
committed by GitHub
parent 1ecc355062
commit 1c3ce17f13
8 changed files with 221 additions and 225 deletions
+27 -5
View File
@@ -170,11 +170,11 @@ __device__ void QueuePair::post_wqe_rma_mt(int pe, int32_t size, uintptr_t *ladd
}
}
__device__ void QueuePair::post_wqe_rma_single(int pe, int32_t size, uintptr_t *laddr, uintptr_t *raddr, uint8_t opcode) {
__device__ void QueuePair::post_wqe_rma_single(int32_t size, uintptr_t *laddr, uintptr_t *raddr, uint8_t opcode, bool ring_db) {
switch (gda_provider_) {
#if defined(GDA_BNXT)
case GDAProvider::BNXT:
return bnxt_post_wqe_rma_single(pe, size, laddr, raddr, opcode);
return bnxt_post_wqe_rma_single(size, laddr, raddr, opcode, ring_db);
#endif
case GDAProvider::IONIC:
case GDAProvider::MLX5:
@@ -192,7 +192,7 @@ __device__ uint64_t QueuePair::post_wqe_amo(int pe, int32_t size, uintptr_t *rad
#endif
#if defined(GDA_BNXT)
case GDAProvider::BNXT:
return bnxt_post_wqe_amo(pe, size, raddr, opcode, atomic_data, atomic_cmp, fetching);
return bnxt_post_wqe_amo(raddr, opcode, atomic_data, atomic_cmp, fetching);
#endif
#if defined(GDA_IONIC)
case GDAProvider::IONIC:
@@ -204,6 +204,22 @@ __device__ uint64_t QueuePair::post_wqe_amo(int pe, int32_t size, uintptr_t *rad
}
}
__device__ uint64_t QueuePair::post_wqe_amo_single(uintptr_t *raddr, uint8_t opcode,
int64_t atomic_data, int64_t atomic_cmp,
bool fetching) {
switch (gda_provider_) {
#if defined(GDA_BNXT)
case GDAProvider::BNXT:
return bnxt_post_wqe_amo_single(raddr, opcode, atomic_data, atomic_cmp, fetching);
#endif
case GDAProvider::MLX5:
case GDAProvider::IONIC:
default:
assert(false /* invalid nic provider */);
return 0;
}
}
__device__ void QueuePair::quiet(Collectivity cy) {
switch (gda_provider_) {
#if defined(GDA_MLX5)
@@ -253,10 +269,10 @@ __device__ void QueuePair::put_nbi(void *dest, const void *source, size_t nelems
post_wqe_rma(pe, nelems, src, dst, gda_op_rdma_write, cy);
}
__device__ void QueuePair::put_nbi_single(void *dest, const void *source, size_t nelems, int pe) {
__device__ void QueuePair::put_nbi_single(void *dest, const void *source, size_t nelems, bool ring_db) {
uintptr_t *src = reinterpret_cast<uintptr_t*>(const_cast<void*>(source));
uintptr_t *dst = reinterpret_cast<uintptr_t*>(dest);
post_wqe_rma_single(pe, nelems, src, dst, gda_op_rdma_write);
post_wqe_rma_single(nelems, src, dst, gda_op_rdma_write, ring_db);
}
__device__ void QueuePair::get_nbi(void *dest, const void *source, size_t nelems, int pe, Collectivity cy) {
@@ -285,4 +301,10 @@ __device__ void QueuePair::atomic_nofetch(void *dest, int64_t atomic_data, int64
post_wqe_amo(pe, sizeof(int64_t), dst, gda_op_atomic_fa, atomic_data, atomic_cmp, false);
}
__device__ void QueuePair::atomic_nofetch_single(void *dest, int64_t value) {
const bool fetching = false;
uintptr_t *dst = static_cast<uintptr_t*>(dest);
post_wqe_amo_single(dst, gda_op_atomic_fa, value, 0, false);
}
} // namespace rocshmem