diff --git a/cmake/modules/ROCM.cmake b/cmake/modules/ROCM.cmake index ce9ad1414f90..7b257d3d526d 100644 --- a/cmake/modules/ROCM.cmake +++ b/cmake/modules/ROCM.cmake @@ -34,11 +34,6 @@ if(USE_ROCM) tvm_file_glob(GLOB RUNTIME_ROCM_SRCS src/backend/rocm/runtime/*.cc) - set(_rocm_libs ${ROCM_HIPHCC_LIBRARY}) - if(ROCM_HSA_LIBRARY) - list(APPEND _rocm_libs ${ROCM_HSA_LIBRARY}) - endif() - add_library(tvm_runtime_rocm_objs OBJECT ${RUNTIME_ROCM_SRCS}) target_link_libraries(tvm_runtime_rocm_objs PUBLIC tvm_ffi_header) set_target_properties(tvm_runtime_rocm_objs PROPERTIES POSITION_INDEPENDENT_CODE ON) @@ -47,7 +42,7 @@ if(USE_ROCM) endif() add_library(tvm_runtime_rocm SHARED $) list(APPEND TVM_RUNTIME_BACKEND_LIBS tvm_runtime_rocm) - target_link_libraries(tvm_runtime_rocm PUBLIC tvm_runtime ${_rocm_libs}) + target_link_libraries(tvm_runtime_rocm PUBLIC tvm_runtime ${ROCM_HIPHCC_LIBRARY}) tvm_configure_target_library(tvm_runtime_rocm RUNTIME_MODULE) endif(USE_ROCM) diff --git a/cmake/utils/FindROCM.cmake b/cmake/utils/FindROCM.cmake index 6bf509cdba77..3c23cd49793b 100644 --- a/cmake/utils/FindROCM.cmake +++ b/cmake/utils/FindROCM.cmake @@ -53,7 +53,6 @@ macro(find_rocm use_rocm) endif() find_library(ROCM_HIPBLAS_LIBRARY hipblas ${__rocm_sdk}/lib) find_library(ROCM_HIPBLASLT_LIBRARY hipblaslt ${__rocm_sdk}/lib) - find_library(ROCM_HSA_LIBRARY hsa-runtime64 ${__rocm_sdk}/lib) if(ROCM_HIPHCC_LIBRARY) set(ROCM_FOUND TRUE) diff --git a/src/backend/rocm/runtime/rocm_device_api.cc b/src/backend/rocm/runtime/rocm_device_api.cc index 6612f1a8fb84..a5c0051848dd 100644 --- a/src/backend/rocm/runtime/rocm_device_api.cc +++ b/src/backend/rocm/runtime/rocm_device_api.cc @@ -22,7 +22,6 @@ * \brief GPU specific API */ #include -#include #include #include #include @@ -42,14 +41,10 @@ class ROCMDeviceAPI final : public DeviceAPI { int value = 0; switch (kind) { case kExist: { - if (hsa_init() == HSA_STATUS_SUCCESS) { - int dev; - ROCM_CALL(hipGetDeviceCount(&dev)); - value = dev > device.device_id ? 1 : 0; - hsa_shut_down(); - } else { - value = 0; - } + // Missing devices or an incompatible driver must return false, not throw. + int count = 0; + hipError_t status = hipGetDeviceCount(&count); + value = status == hipSuccess && device.device_id >= 0 && device.device_id < count; break; } case kMaxThreadsPerBlock: { diff --git a/tests/python/runtime/test_runtime_device_api.py b/tests/python/runtime/test_runtime_device_api.py index 8c4ec430f1da..6008cc9f33c9 100644 --- a/tests/python/runtime/test_runtime_device_api.py +++ b/tests/python/runtime/test_runtime_device_api.py @@ -19,6 +19,8 @@ import subprocess import sys +import pytest + import tvm import tvm.testing @@ -48,5 +50,30 @@ def test_check_if_device_exists(): ) +@pytest.mark.skipif( + not tvm.runtime.enabled("rocm"), + reason="Requires the ROCm runtime to be built", +) +@pytest.mark.parametrize("device_id", [-1, 2**31 - 1]) +def test_rocm_invalid_device_does_not_exist(device_id): + assert not tvm.rocm(device_id).exist + + +@pytest.mark.skipif( + not tvm.runtime.enabled("rocm"), + reason="Requires the ROCm runtime to be built", +) +def test_rocm_hidden_device_does_not_exist(): + subprocess.check_call( + [sys.executable, "-c", "import tvm; assert not tvm.rocm(0).exist"], + env={**os.environ, "HIP_VISIBLE_DEVICES": "", "ROCR_VISIBLE_DEVICES": ""}, + ) + + +@tvm.testing.requires_rocm +def test_rocm_device_exists(): + assert tvm.rocm(0).exist + + if __name__ == "__main__": tvm.testing.main()