diff --git a/hsa.cpp b/hsa.cpp index c6fe8b2e9a..431e7bb91a 100755 --- a/hsa.cpp +++ b/hsa.cpp @@ -2,23 +2,13 @@ #include "impl/hsa/hsa.h" #include "impl/hsa/hsa_ven_amd_loader.h" -static std::unique_ptr lock_ = std::make_unique(); -static hsa_status_t (*fn_hsa_ven_amd_loader_query_host_address)( - const void *device_address, const void **host_address); - -#if 0 -static hsa_signal_value_t (*fn_hsa_signal_load_relaxed)(hsa_signal_t signal); -static hsa_signal_value_t (*fn_hsa_signal_wait_relaxed)( - hsa_signal_t signal, hsa_signal_condition_t condition, - hsa_signal_value_t compare_value, uint64_t timeout_hint, - hsa_wait_state_t wait_state_hint); -static void (*fn_hsa_signal_store_screlease)(hsa_signal_t hsa_signal, - hsa_signal_value_t value); +static std::mutex* lock_ = new std::mutex(); +#if 1 #define _HSAKMT_LOOKUP_SYMS(_sym) \ -if (_sym == nullptr) { \ +if (fn_##_sym == nullptr) { \ std::lock_guard gard(*lock_); \ - if (_sym == nullptr) { \ + if (fn_##_sym == nullptr) { \ fn_##_sym = \ reinterpret_cast(dlsym(RTLD_DEFAULT, #_sym)); \ if (!fn_##_sym) { \ @@ -34,9 +24,22 @@ do { \ } \ } while(0); +bool hsakmt_hsa_loader_init() { + void *hsa_loader_handle = dlopen("libhsa-runtime64.so", RTLD_NOW | RTLD_GLOBAL); + if (hsa_loader_handle == nullptr) { + pr_err("dlopen libhsa-runtime64.so failed - %s\n", dlerror()); + return false; + } + dlclose(hsa_loader_handle); + return true; +} + hsa_signal_value_t hsakmt_hsa_signal_load_relaxed(hsa_signal_t signal) { + static hsa_signal_value_t (*fn_hsa_signal_load_relaxed)(hsa_signal_t signal) = nullptr; + _HSAKMT_LOOKUP_SYMS(hsa_signal_load_relaxed); _HSAKMT_EXEC_API(hsa_signal_load_relaxed, signal); + return 0; } @@ -44,20 +47,32 @@ hsa_signal_value_t hsakmt_hsa_signal_wait_relaxed( hsa_signal_t signal, hsa_signal_condition_t condition, hsa_signal_value_t compare_value, uint64_t timeout_hint, hsa_wait_state_t wait_state_hint) { +static hsa_signal_value_t (*fn_hsa_signal_wait_relaxed)( + hsa_signal_t signal, hsa_signal_condition_t condition, + hsa_signal_value_t compare_value, uint64_t timeout_hint, + hsa_wait_state_t wait_state_hint) = nullptr; + _HSAKMT_LOOKUP_SYMS(hsa_signal_wait_relaxed); _HSAKMT_EXEC_API(hsa_signal_wait_relaxed, signal, condition, compare_value, timeout_hint, wait_state_hint); + return 0; } void hsakmt_hsa_signal_store_screlease(hsa_signal_t hsa_signal, hsa_signal_value_t value){ +static void (*fn_hsa_signal_store_screlease)(hsa_signal_t hsa_signal, + hsa_signal_value_t value) = nullptr; + _HSAKMT_LOOKUP_SYMS(hsa_signal_store_screlease); _HSAKMT_EXEC_API(hsa_signal_store_screlease, hsa_signal, value); } hsa_status_t hsakmt_hsa_ven_amd_loader_query_host_address( const void *device_address, const void **host_address) { + static hsa_status_t (*fn_hsa_ven_amd_loader_query_host_address)( + const void *device_address, const void **host_address) = nullptr; + if (fn_hsa_ven_amd_loader_query_host_address == nullptr) { std::lock_guard gard(*lock_); if (fn_hsa_ven_amd_loader_query_host_address == nullptr) { @@ -101,6 +116,9 @@ void hsakmt_hsa_signal_store_screlease(hsa_signal_t hsa_signal, hsa_status_t hsakmt_hsa_ven_amd_loader_query_host_address( const void *device_address, const void **host_address) { + static hsa_status_t (*fn_hsa_ven_amd_loader_query_host_address)( + const void *device_address, const void **host_address) = nullptr; + if (fn_hsa_ven_amd_loader_query_host_address == nullptr) { std::lock_guard gard(*lock_); if (fn_hsa_ven_amd_loader_query_host_address == nullptr) { diff --git a/librocdxg.h b/librocdxg.h index a36fba1ff7..0cdbdd2c53 100644 --- a/librocdxg.h +++ b/librocdxg.h @@ -284,4 +284,5 @@ HSAKMT_STATUS import_dmabuf_fd(int DMABufFd, bool is_ipc_memfd, wsl::thunk::GpuMemoryHandle *GpuMemHandle); +bool hsakmt_hsa_loader_init(); #endif diff --git a/openclose.cpp b/openclose.cpp index 71a9977e31..6be1b24b86 100644 --- a/openclose.cpp +++ b/openclose.cpp @@ -557,6 +557,7 @@ HSAKMT_STATUS HSAKMTAPI hsaKmtOpenKFD(void) { dxg_runtime->dxg_fd = fd; } + hsakmt_hsa_loader_init(); init_page_size(); char *useSvmStr = getenv("HSA_USE_SVM");