From d81343a0611fa404c54f11a1b1dc78c77662f7b1 Mon Sep 17 00:00:00 2001 From: Yichen Yan Date: Sun, 4 Oct 2026 13:38:55 +0800 Subject: [PATCH 1/3] [FIX][ROCm] Query Windows device existence through HIP --- src/backend/rocm/runtime/rocm_device_api.cc | 11 ++++++++++ .../python/runtime/test_runtime_device_api.py | 22 +++++++++++++++++++ 2 files changed, 33 insertions(+) diff --git a/src/backend/rocm/runtime/rocm_device_api.cc b/src/backend/rocm/runtime/rocm_device_api.cc index 6612f1a8fb84..7a62fd537ac6 100644 --- a/src/backend/rocm/runtime/rocm_device_api.cc +++ b/src/backend/rocm/runtime/rocm_device_api.cc @@ -22,7 +22,9 @@ * \brief GPU specific API */ #include +#ifndef _WIN32 #include +#endif #include #include #include @@ -42,6 +44,14 @@ class ROCMDeviceAPI final : public DeviceAPI { int value = 0; switch (kind) { case kExist: { +#ifdef _WIN32 + // Windows HIP has no HSA runtime. Missing devices or an incompatible + // driver must make an existence query return false, not throw. + int dev = 0; + if (hipGetDeviceCount(&dev) == hipSuccess) { + value = device.device_id >= 0 && device.device_id < dev; + } +#else if (hsa_init() == HSA_STATUS_SUCCESS) { int dev; ROCM_CALL(hipGetDeviceCount(&dev)); @@ -50,6 +60,7 @@ class ROCMDeviceAPI final : public DeviceAPI { } else { value = 0; } +#endif 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..489189d50071 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,25 @@ def test_check_if_device_exists(): ) +@pytest.mark.skipif( + sys.platform != "win32" or not tvm.runtime.enabled("rocm"), + reason="Requires the Windows HIP runtime", +) +def test_windows_rocm_invalid_device_does_not_exist(): + assert not tvm.rocm(-1).exist + assert not tvm.rocm(2**31 - 1).exist + + +@pytest.mark.skipif( + sys.platform != "win32" or not tvm.runtime.enabled("rocm"), + reason="Requires the Windows HIP runtime", +) +def test_windows_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": ""}, + ) + + if __name__ == "__main__": tvm.testing.main() From ee83b20b7de9dc6564b49092ce86839a189bab57 Mon Sep 17 00:00:00 2001 From: Yichen Yan Date: Sun, 4 Oct 2026 13:50:53 +0800 Subject: [PATCH 2/3] [DOCS][ROCm] Clarify Windows SDK header availability --- src/backend/rocm/runtime/rocm_device_api.cc | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/backend/rocm/runtime/rocm_device_api.cc b/src/backend/rocm/runtime/rocm_device_api.cc index 7a62fd537ac6..914eb7dcf7e7 100644 --- a/src/backend/rocm/runtime/rocm_device_api.cc +++ b/src/backend/rocm/runtime/rocm_device_api.cc @@ -45,8 +45,8 @@ class ROCMDeviceAPI final : public DeviceAPI { switch (kind) { case kExist: { #ifdef _WIN32 - // Windows HIP has no HSA runtime. Missing devices or an incompatible - // driver must make an existence query return false, not throw. + // Windows HIP SDK packages can omit the HSA development headers. + // Missing devices or an incompatible driver must return false, not throw. int dev = 0; if (hipGetDeviceCount(&dev) == hipSuccess) { value = device.device_id >= 0 && device.device_id < dev; From d3eddcad6ab7debf825911a8f19a4743bc2549e1 Mon Sep 17 00:00:00 2001 From: Yichen Yan Date: Sun, 4 Oct 2026 14:07:36 +0800 Subject: [PATCH 3/3] [FIX][ROCm] Use HIP device queries across platforms --- cmake/modules/ROCM.cmake | 7 +----- cmake/utils/FindROCM.cmake | 1 - src/backend/rocm/runtime/rocm_device_api.cc | 22 +++--------------- .../python/runtime/test_runtime_device_api.py | 23 +++++++++++-------- 4 files changed, 18 insertions(+), 35 deletions(-) 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 914eb7dcf7e7..a5c0051848dd 100644 --- a/src/backend/rocm/runtime/rocm_device_api.cc +++ b/src/backend/rocm/runtime/rocm_device_api.cc @@ -22,9 +22,6 @@ * \brief GPU specific API */ #include -#ifndef _WIN32 -#include -#endif #include #include #include @@ -44,23 +41,10 @@ class ROCMDeviceAPI final : public DeviceAPI { int value = 0; switch (kind) { case kExist: { -#ifdef _WIN32 - // Windows HIP SDK packages can omit the HSA development headers. // Missing devices or an incompatible driver must return false, not throw. - int dev = 0; - if (hipGetDeviceCount(&dev) == hipSuccess) { - value = device.device_id >= 0 && device.device_id < dev; - } -#else - 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; - } -#endif + 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 489189d50071..6008cc9f33c9 100644 --- a/tests/python/runtime/test_runtime_device_api.py +++ b/tests/python/runtime/test_runtime_device_api.py @@ -51,24 +51,29 @@ def test_check_if_device_exists(): @pytest.mark.skipif( - sys.platform != "win32" or not tvm.runtime.enabled("rocm"), - reason="Requires the Windows HIP runtime", + not tvm.runtime.enabled("rocm"), + reason="Requires the ROCm runtime to be built", ) -def test_windows_rocm_invalid_device_does_not_exist(): - assert not tvm.rocm(-1).exist - assert not tvm.rocm(2**31 - 1).exist +@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( - sys.platform != "win32" or not tvm.runtime.enabled("rocm"), - reason="Requires the Windows HIP runtime", + not tvm.runtime.enabled("rocm"), + reason="Requires the ROCm runtime to be built", ) -def test_windows_rocm_hidden_device_does_not_exist(): +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": ""}, + 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()