diff --git a/projects/rccl/src/misc/rocm_smi_wrap.cc b/projects/rccl/src/misc/rocm_smi_wrap.cc index 75085f2a1b..3589f41906 100644 --- a/projects/rccl/src/misc/rocm_smi_wrap.cc +++ b/projects/rccl/src/misc/rocm_smi_wrap.cc @@ -47,9 +47,9 @@ static int is_wsl2 = -1; RCCL_PARAM(UseRocmSmiLib, "USE_ROCM_SMI_LIB", 0); // Opt-in environment variable for enabling using rocm_smi_lib instead of internal code ncclResult_t rocm_smi_init() { - if (is_wsl2 == -1) - is_wsl2 = (access("/dev/dxg", F_OK) == -1) ? 0 : 1; - if (is_wsl2) { + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE) == -1) + __atomic_store_n(&is_wsl2, (access("/dev/dxg", F_OK) == -1) ? 0 : 1, __ATOMIC_RELEASE); + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE)) { INFO(NCCL_INIT, "Not using rocm_smi_lib due to WSL2 environment detected."); return ncclSuccess; } @@ -71,7 +71,7 @@ ncclResult_t rocm_smi_init() { } ncclResult_t rocm_smi_getNumDevice(uint32_t* num_devs) { - if (is_wsl2) + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE)) CUDACHECK(cudaGetDeviceCount((int *)num_devs)); else if (rcclParamUseRocmSmiLib()) { @@ -85,7 +85,7 @@ ncclResult_t rocm_smi_getNumDevice(uint32_t* num_devs) { ncclResult_t rocm_smi_getDevicePciBusIdString(uint32_t deviceIndex, char* busId, size_t len) { uint64_t id; - if (is_wsl2) { + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE)) { CUDACHECK(cudaDeviceGetPCIBusId(busId, len, deviceIndex)); } else { /** rocm_smi's bus ID format @@ -109,7 +109,7 @@ ncclResult_t rocm_smi_getDevicePciBusIdString(uint32_t deviceIndex, char* busId, ncclResult_t rocm_smi_getDeviceIndexByPciBusId(const char* pciBusId, uint32_t* deviceIndex) { - if (is_wsl2) { + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE)) { CUDACHECK(hipDeviceGetByPCIBusId((int *)deviceIndex, pciBusId)); return ncclSuccess; } else { @@ -160,7 +160,7 @@ ncclResult_t rocm_smi_getDeviceIndexByPciBusId(const char* pciBusId, uint32_t* d } ncclResult_t rocm_smi_getLinkInfo(int srcIndex, int dstIndex, RSMI_IO_LINK_TYPE* rsmi_type, int *hops, int *count) { - if (is_wsl2) { + if (__atomic_load_n(&is_wsl2, __ATOMIC_ACQUIRE)) { *rsmi_type = RSMI_IOLINK_TYPE_PCIEXPRESS; *hops = 1; *count = 1;