From 178c1774354d073e1b9ba919247f1ff7ebbdf192 Mon Sep 17 00:00:00 2001 From: Fred Heinecke Date: Fri, 24 Jul 2026 10:35:06 -0500 Subject: [PATCH 1/2] Enable runtime resolution of CUDA header path for NVRTC Signed-off-by: Fred Heinecke --- transformer_engine/common/CMakeLists.txt | 1 + .../common/util/cuda_runtime.cpp | 78 +++++++++++++++++++ 2 files changed, 79 insertions(+) diff --git a/transformer_engine/common/CMakeLists.txt b/transformer_engine/common/CMakeLists.txt index 8eba515e5e..4b9c34ef6b 100644 --- a/transformer_engine/common/CMakeLists.txt +++ b/transformer_engine/common/CMakeLists.txt @@ -360,6 +360,7 @@ target_link_libraries(transformer_engine PUBLIC CUDA::cublas CUDA::cudart CUDNN::cudnn_all) +target_link_libraries(transformer_engine PRIVATE ${CMAKE_DL_LIBS}) target_include_directories(transformer_engine PRIVATE ${CMAKE_CUDA_TOOLKIT_INCLUDE_DIRECTORIES}) diff --git a/transformer_engine/common/util/cuda_runtime.cpp b/transformer_engine/common/util/cuda_runtime.cpp index 504d761bb1..c811dbdcd0 100644 --- a/transformer_engine/common/util/cuda_runtime.cpp +++ b/transformer_engine/common/util/cuda_runtime.cpp @@ -8,6 +8,8 @@ #include +#include + #include #include #include @@ -26,6 +28,81 @@ namespace { // String with build-time CUDA include path #include "string_path_cuda_include.h" +// Get the runtime directory of the shared library that contains this code +std::filesystem::path shared_library_directory() { + static const char library_anchor = 0; + Dl_info library_info{}; + if (dladdr(static_cast(&library_anchor), &library_info) == 0 || + library_info.dli_fname == nullptr) { + return {}; + } + + std::filesystem::path library_path = library_info.dli_fname; + if (library_path.is_relative()) { + std::error_code error; + library_path = std::filesystem::absolute(library_path, error); + if (error) { + return {}; + } + } + + return library_path.parent_path(); +} + +std::string runtime_cuda_major_version() { + int runtime_version = 0; + // Header discovery is best-effort, so do not throw if the runtime cannot + // report its version. + if (cudaRuntimeGetVersion(&runtime_version) != cudaSuccess || runtime_version <= 0) { + return {}; + } + + return std::to_string(runtime_version / 1000); +} + +std::filesystem::path python_cuda_directory() { + using Path = std::filesystem::path; + + // Find the packages directory by traversing up the directory tree until a known package + // directory is found. + Path directory = shared_library_directory(); + while (true) { + const auto filename = directory.filename(); + if (filename == "site-packages" || filename == "dist-packages") { + break; + } + + const Path parent = directory.parent_path(); + if (parent == directory) { + // Root directory reached + return {}; + } + + directory = parent; + } + + const Path nvidia_directory = directory / "nvidia"; + const auto cuda_major_version = runtime_cuda_major_version(); + if (cuda_major_version.empty()) { + return {}; + } + + std::error_code error; + const Path cuda_directory = nvidia_directory / ("cu" + cuda_major_version); + if (std::filesystem::is_directory(cuda_directory, error)) { + return cuda_directory; + } + + // CUDA 12 Python wheels use the older nvidia/cuda_runtime layout. + error.clear(); + const Path legacy_cuda_directory = nvidia_directory / "cuda_runtime"; + if (std::filesystem::is_directory(legacy_cuda_directory, error)) { + return legacy_cuda_directory; + } + + return {}; +} + } // namespace int num_devices() { @@ -152,6 +229,7 @@ const std::string &include_directory(bool required) { std::vector> search_paths = {{"NVTE_CUDA_INCLUDE_DIR", ""}, {"CUDA_HOME", ""}, {"CUDA_DIR", ""}, + {"", python_cuda_directory()}, {"", string_path_cuda_include}, {"", "/usr/local/cuda"}}; for (auto &[env, p] : search_paths) { From 4c21479126f6d3c52f787e403fab5e2d28602d1e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 24 Jul 2026 15:43:50 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/common/util/cuda_runtime.cpp | 1 - 1 file changed, 1 deletion(-) diff --git a/transformer_engine/common/util/cuda_runtime.cpp b/transformer_engine/common/util/cuda_runtime.cpp index c811dbdcd0..5674feadef 100644 --- a/transformer_engine/common/util/cuda_runtime.cpp +++ b/transformer_engine/common/util/cuda_runtime.cpp @@ -7,7 +7,6 @@ #include "../util/cuda_runtime.h" #include - #include #include