Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 1 addition & 6 deletions cmake/modules/ROCM.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -47,7 +42,7 @@ if(USE_ROCM)
endif()
add_library(tvm_runtime_rocm SHARED $<TARGET_OBJECTS:tvm_runtime_rocm_objs>)
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)

Expand Down
1 change: 0 additions & 1 deletion cmake/utils/FindROCM.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
13 changes: 4 additions & 9 deletions src/backend/rocm/runtime/rocm_device_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@
* \brief GPU specific API
*/
#include <hip/hip_runtime_api.h>
#include <hsa/hsa.h>
#include <tvm/ffi/extra/c_env_api.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
Expand All @@ -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: {
Expand Down
27 changes: 27 additions & 0 deletions tests/python/runtime/test_runtime_device_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
import subprocess
import sys

import pytest

import tvm
import tvm.testing

Expand Down Expand Up @@ -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()
Loading