diff --git a/llvmpy/gen/binding.py b/llvmpy/gen/binding.py index 0f5c4fd..ab8e7aa 100644 --- a/llvmpy/gen/binding.py +++ b/llvmpy/gen/binding.py @@ -838,3 +838,17 @@ class Attr(object): writer.println() +# +# Pick-up environ var +# +PTX_SUPPORT = os.environ['LLVMPY_PTX_SUPPORT'] + +def _parse_llvm_version(ver): + import re + m = re.compile(r'(\d+)\.(\d+)').match(ver) + assert m + major, minor = m.groups() + return int(major), int(minor) + +LLVM_VERSION = _parse_llvm_version(os.environ['LLVMPY_LLVM_VERSION']) + diff --git a/llvmpy/src/Support/TargetSelect.py b/llvmpy/src/Support/TargetSelect.py index ca97c24..85e1d74 100644 --- a/llvmpy/src/Support/TargetSelect.py +++ b/llvmpy/src/Support/TargetSelect.py @@ -17,12 +17,14 @@ InitializeNativeTargetDisassembler = llvm.Function( #InitializeAllTargetMCs = llvm.Function('InitializeAllTargetMCs') #InitializeAllAsmPrinters = llvm.Function('InitializeAllAsmPrinters') -#LLVMInitializePTXTarget = default.Function('LLVMInitializePTXTarget') -#LLVMInitializePTXTargetInfo = default.Function('LLVMInitializePTXTargetInfo') -#LLVMInitializePTXTargetMC = default.Function('LLVMInitializePTXTargetMC') -#LLVMInitializePTXAsmPrinter = default.Function('LLVMInitializePTXAsmPrinter') +if PTX_SUPPORT == 'PTX': + LLVMInitializePTXTarget = default.Function('LLVMInitializePTXTarget') + LLVMInitializePTXTargetInfo = default.Function('LLVMInitializePTXTargetInfo') + LLVMInitializePTXTargetMC = default.Function('LLVMInitializePTXTargetMC') + LLVMInitializePTXAsmPrinter = default.Function('LLVMInitializePTXAsmPrinter') -LLVMInitializeNVPTXTarget = default.Function('LLVMInitializeNVPTXTarget') -LLVMInitializeNVPTXTargetInfo = default.Function('LLVMInitializeNVPTXTargetInfo') -LLVMInitializeNVPTXTargetMC = default.Function('LLVMInitializeNVPTXTargetMC') -LLVMInitializeNVPTXAsmPrinter = default.Function('LLVMInitializeNVPTXAsmPrinter') +if PTX_SUPPORT == 'NVPTX': + LLVMInitializeNVPTXTarget = default.Function('LLVMInitializeNVPTXTarget') + LLVMInitializeNVPTXTargetInfo = default.Function('LLVMInitializeNVPTXTargetInfo') + LLVMInitializeNVPTXTargetMC = default.Function('LLVMInitializeNVPTXTargetMC') + LLVMInitializeNVPTXAsmPrinter = default.Function('LLVMInitializeNVPTXAsmPrinter') diff --git a/setup.py b/setup.py index 9bcd9f0..6fa6653 100644 --- a/setup.py +++ b/setup.py @@ -97,7 +97,7 @@ def determine_to_use_dynlink(libdir, llvm_version): dynlink = determine_to_use_dynlink(libdir, llvm_version) - +ptx_support = '' if dynlink: print('Using dynamic linking') @@ -120,12 +120,13 @@ else: if (nvptx_components & enabled_components) == nvptx_components: print("Using NVPTX") extra_components.extend(nvptx_components) + ptx_support = 'NVPTX' elif (ptx_components & enabled_components) == ptx_components: print("Using PTX") extra_components.extend(ptx_components) + ptx_support = 'PTX' else: print("No CUDA support") - macros.append(('LLVM_DISABLE_PTX', None)) libs_core, objs_core = get_libs_and_objs( ['core', 'analysis', 'scalaropts', 'executionengine', @@ -136,14 +137,17 @@ else: if sys.platform == 'win32': # If no PTX lib got added, disable PTX in the build - if 'LLVMPTXCodeGen' not in libs_core: - macros.append(('LLVM_DISABLE_PTX', None)) + if 'LLVMPTXCodeGen' in libs_core: + ptx_support = 'ptx' else: macros.append(('_GNU_SOURCE', None)) # auto generate bindings +os.environ['LLVMPY_PTX_SUPPORT'] = ptx_support +os.environ['LLVMPY_LLVM_VERSION'] = llvm_version check_call([sys.executable, 'llvmpy/build.py']) +# generate shared objects extra_link_args = ldflags.split() kwds = dict( ext_modules = [