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:
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user