codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import ctypes
6import sys
7import warnings
8
9
10def find_cudart_versions(build_env=False, build_cuda_version=None):
11 # ctypes.CDLL and ctypes.util.find_library load the latest installed library.
12 # it may not the the library that would be loaded by onnxruntime.
13 # for example, in an environment with Cuda 11.1 and subsequently
14 # conda cudatoolkit 10.2.89 installed. ctypes will find cudart 10.2. however,
15 # onnxruntime built with Cuda 11.1 will find and load cudart for Cuda 11.1.
16 # for the above reason, we need find all versions in the environment and
17 # only give warnings if the expected cuda version is not found.
18 # in onnxruntime build environment, we expected only one Cuda version.
19 if not sys.platform.startswith("linux"):
20 warnings.warn("find_cudart_versions only works on Linux")
21 return None
22
23 cudart_possible_versions = {None, build_cuda_version}
24
25 def get_cudart_version(find_cudart_version=None):
26 cudart_lib_filename = "libcudart.so"
27 if find_cudart_version:
28 cudart_lib_filename = cudart_lib_filename + "." + find_cudart_version
29
30 try:
31 cudart = ctypes.CDLL(cudart_lib_filename)
32 cudart.cudaRuntimeGetVersion.restype = int
33 cudart.cudaRuntimeGetVersion.argtypes = [ctypes.POINTER(ctypes.c_int)]
34 version = ctypes.c_int()
35 status = cudart.cudaRuntimeGetVersion(ctypes.byref(version))
36 if status != 0:
37 return None
38 except Exception:
39 return None
40
41 return version.value
42
43 # use set to avoid duplications
44 cudart_found_versions = {get_cudart_version(cudart_version) for cudart_version in cudart_possible_versions}
45
46 # convert to list and remove None
47 return [ver for ver in cudart_found_versions if ver]
48 