Make WSL detection thread safe (#1178)

* Make WSL detection thread safe

* Change static to beginning

* Switch to use atomics

[ROCm/rccl commit: 8f099b1adb]
This commit is contained in:
Wenkai Du
2024-05-28 17:23:50 -07:00
committed by GitHub
parent 7f4bec7682
commit ae6c372406
+7 -7
View File
@@ -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;