microsoft/onnxruntime-extensions
Publicmirrored from https://github.com/microsoft/onnxruntime-extensionsAvailable
cmake/ext_cuda.cmake
87lines · modecode
| 1 | # Copyright (c) Microsoft Corporation. All rights reserved. |
| 2 | # Licensed under the MIT License. |
| 3 | |
| 4 | enable_language(CUDA) |
| 5 | |
| 6 | set(CMAKE_CUDA_RUNTIME_LIBRARY Shared) |
| 7 | set(CMAKE_CUDA_STANDARD 17) |
| 8 | include(CMakeDependentOption) |
| 9 | cmake_dependent_option(USE_FLASH_ATTENTION "Build flash attention kernel for scaled dot product attention" ON "NOT WIN32" OFF) |
| 10 | option(USE_MEMORY_EFFICIENT_ATTENTION "Build memory efficient attention kernel for scaled dot product attention" ON) |
| 11 | if (CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 11.6) |
| 12 | message( STATUS "Turn off flash attention and memory efficient attention since CUDA compiler version < 11.6") |
| 13 | set(USE_FLASH_ATTENTION OFF) |
| 14 | set(USE_MEMORY_EFFICIENT_ATTENTION OFF) |
| 15 | endif() |
| 16 | |
| 17 | |
| 18 | if(NOT CMAKE_CUDA_ARCHITECTURES) |
| 19 | if(CMAKE_LIBRARY_ARCHITECTURE STREQUAL "aarch64-linux-gnu") |
| 20 | # Support for Jetson/Tegra ARM devices |
| 21 | set(CMAKE_CUDA_FLAGS |
| 22 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_53,code=sm_53") # TX1, Nano |
| 23 | set(CMAKE_CUDA_FLAGS |
| 24 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_62,code=sm_62") # TX2 |
| 25 | set(CMAKE_CUDA_FLAGS |
| 26 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_72,code=sm_72") # AGX Xavier, |
| 27 | # NX Xavier |
| 28 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 11) |
| 29 | set(CMAKE_CUDA_FLAGS |
| 30 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_87,code=sm_87") # AGX Orin, |
| 31 | # NX Orin |
| 32 | endif() |
| 33 | else() |
| 34 | # the following compute capabilities are removed in CUDA 11 Toolkit |
| 35 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 11) |
| 36 | set(CMAKE_CUDA_FLAGS |
| 37 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_30,code=sm_30") # K series |
| 38 | endif() |
| 39 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_LESS 12) |
| 40 | # 37, 50 still work in CUDA 11 but are marked deprecated and will be |
| 41 | # removed in future CUDA version. |
| 42 | set(CMAKE_CUDA_FLAGS |
| 43 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_37,code=sm_37") # K80 |
| 44 | set(CMAKE_CUDA_FLAGS |
| 45 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_50,code=sm_50") # M series |
| 46 | endif() |
| 47 | set(CMAKE_CUDA_FLAGS |
| 48 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_52,code=sm_52") # M60 |
| 49 | set(CMAKE_CUDA_FLAGS |
| 50 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_60,code=sm_60") # P series |
| 51 | set(CMAKE_CUDA_FLAGS |
| 52 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_70,code=sm_70") # V series |
| 53 | set(CMAKE_CUDA_FLAGS |
| 54 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_75,code=sm_75") # T series |
| 55 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 11) |
| 56 | set(CMAKE_CUDA_FLAGS |
| 57 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_80,code=sm_80") # A series |
| 58 | endif() |
| 59 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 12) |
| 60 | set(CMAKE_CUDA_FLAGS |
| 61 | "${CMAKE_CUDA_FLAGS} -gencode=arch=compute_90,code=sm_90") # H series |
| 62 | endif() |
| 63 | endif() |
| 64 | endif() |
| 65 | set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} --expt-relaxed-constexpr") |
| 66 | if(CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL 11) |
| 67 | set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} --Werror default-stream-launch") |
| 68 | endif() |
| 69 | |
| 70 | if(NOT WIN32) |
| 71 | list(APPEND CUDA_NVCC_FLAGS --compiler-options -fPIC) |
| 72 | endif() |
| 73 | |
| 74 | # Options passed to cudafe |
| 75 | set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -Xcudafe \"--diag_suppress=bad_friend_decl\"") |
| 76 | set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -Xcudafe \"--diag_suppress=unsigned_compare_with_zero\"") |
| 77 | set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} -Xcudafe \"--diag_suppress=expr_has_no_effect\"") |
| 78 | |
| 79 | add_compile_definitions(USE_CUDA) |
| 80 | if (USE_FLASH_ATTENTION) |
| 81 | message( STATUS "Enable flash attention") |
| 82 | add_compile_definitions(USE_FLASH_ATTENTION) |
| 83 | endif() |
| 84 | if (USE_MEMORY_EFFICIENT_ATTENTION) |
| 85 | message( STATUS "Enable memory efficient attention") |
| 86 | add_compile_definitions(USE_MEMORY_EFFICIENT_ATTENTION) |
| 87 | endif() |