diff --git a/CMakeLists.txt b/CMakeLists.txt index 3433480..d8dfc00 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -13,7 +13,7 @@ if(CMAKE_CUDA_COMPILER) add_definitions(-DWITH_CUDA) find_package(CUDAToolkit REQUIRED) - set(CMAKE_CUDA_STANDARD 17) + set(CMAKE_CUDA_STANDARD 20) set(CMAKE_CUDA_STANDARD_REQUIRED ON) message(STATUS "INSTALLING EXTENSIONS WITH CUDA!") @@ -25,7 +25,8 @@ else() message(WARNING "NO CUDA INSTALLATION FOUND, TRYING TO INSTALL CPU VERSION ONLY!") endif() -set(CMAKE_CXX_STANDARD 17) +# torch >= 2.6 headers #error without C++20 (torch/all.h, ATen.h). +set(CMAKE_CXX_STANDARD 20) set(CMAKE_CXX_STANDARD_REQUIRED ON) set(CMAKE_POSITION_INDEPENDENT_CODE ON) @@ -40,6 +41,22 @@ find_package(Python COMPONENTS Interpreter Development.Module REQUIRED) find_package(pybind11 CONFIG REQUIRED) find_package(OpenMP) +# Ask torch itself where its CMake config lives: on RHEL-family platforms the +# build env splits site-packages into lib/ and lib64/, and a bare find_package +# never looks inside either. +# AUTOLOAD=0: a torch backend extension (e.g. torch-rbln) must not have to +# load just to answer a path query. +execute_process( + COMMAND "${CMAKE_COMMAND}" -E env TORCH_DEVICE_BACKEND_AUTOLOAD=0 + "${Python_EXECUTABLE}" -c "import torch; print(torch.utils.cmake_prefix_path)" + OUTPUT_VARIABLE TORCH_CMAKE_PREFIX_PATH + OUTPUT_STRIP_TRAILING_WHITESPACE + RESULT_VARIABLE TORCH_QUERY_RESULT +) +if(TORCH_QUERY_RESULT EQUAL 0) + list(APPEND CMAKE_PREFIX_PATH "${TORCH_CMAKE_PREFIX_PATH}") +endif() + find_package(Torch CONFIG REQUIRED) find_library(TORCH_PYTHON_LIBRARY torch_python PATH "${TORCH_INSTALL_PREFIX}/lib") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} ${TORCH_CXX_FLAGS}") diff --git a/pyproject.toml b/pyproject.toml index 4a80e71..11e0694 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,7 +36,7 @@ license-files = ["LICENSE"] exclude = ["**/.mypy_cache/**", "**/build/**", "**/.vscode/**"] [build-system] -requires = ["scikit-build-core>=0.10", "pybind11>=2.10", "cmake", "ninja"] +requires = ["scikit-build-core>=0.10", "pybind11>=2.10", "cmake", "ninja", "torch"] build-backend = "scikit_build_core.build" [tool.isort]