add NVPTX for LLVM 3.2;
auto-gen _intrinsic_ids.py in setup.py;
This commit is contained in:
parent
505d280fe0
commit
d1bcdd05d1
9 changed files with 99 additions and 1961 deletions
|
|
@ -43,6 +43,11 @@
|
||||||
#include "llvm-c/ExecutionEngine.h"
|
#include "llvm-c/ExecutionEngine.h"
|
||||||
#include "llvm-c/Target.h"
|
#include "llvm-c/Target.h"
|
||||||
#include "llvm-c/Transforms/IPO.h"
|
#include "llvm-c/Transforms/IPO.h"
|
||||||
|
#if LLVM_VERSION_MAJOR >= 3 && LLVM_VERSION_MINOR >= 2
|
||||||
|
#include "llvm-c/Linker.h"
|
||||||
|
#else
|
||||||
|
typedef unsigned int LLVMLinkerMode;
|
||||||
|
#endif
|
||||||
|
|
||||||
/* libc includes */
|
/* libc includes */
|
||||||
#include <stdarg.h> /* for malloc(), free() */
|
#include <stdarg.h> /* for malloc(), free() */
|
||||||
|
|
@ -225,7 +230,7 @@ _wLLVMLinkModules(PyObject *self, PyObject *args)
|
||||||
dest = (LLVMModuleRef) PyCapsule_GetPointer(dest_obj, NULL);
|
dest = (LLVMModuleRef) PyCapsule_GetPointer(dest_obj, NULL);
|
||||||
src = (LLVMModuleRef) PyCapsule_GetPointer(src_obj, NULL);
|
src = (LLVMModuleRef) PyCapsule_GetPointer(src_obj, NULL);
|
||||||
|
|
||||||
if (!LLVMLinkModules(dest, src, mode, &errmsg)) {
|
if (!LLVMLinkModules(dest, src, (LLVMLinkerMode)mode, &errmsg)) {
|
||||||
if (errmsg) {
|
if (errmsg) {
|
||||||
ret = PyUnicode_FromString(errmsg);
|
ret = PyUnicode_FromString(errmsg);
|
||||||
LLVMDisposeMessage(errmsg);
|
LLVMDisposeMessage(errmsg);
|
||||||
|
|
@ -897,11 +902,17 @@ _wrap_none2none(LLVMInitializePasses)
|
||||||
|
|
||||||
_wrap_none2obj(LLVMInitializeNativeTarget, int)
|
_wrap_none2obj(LLVMInitializeNativeTarget, int)
|
||||||
_wrap_none2obj(LLVMInitializeNativeTargetAsmPrinter, int)
|
_wrap_none2obj(LLVMInitializeNativeTargetAsmPrinter, int)
|
||||||
|
#if LLVM_HAS_NVPTX
|
||||||
|
_wrap_none2none(LLVMInitializeNVPTXTarget)
|
||||||
|
_wrap_none2none(LLVMInitializeNVPTXTargetInfo)
|
||||||
|
_wrap_none2none( LLVMInitializeNVPTXTargetMC )
|
||||||
|
_wrap_none2none(LLVMInitializeNVPTXAsmPrinter)
|
||||||
|
#else
|
||||||
_wrap_none2none(LLVMInitializePTXTarget)
|
_wrap_none2none(LLVMInitializePTXTarget)
|
||||||
_wrap_none2none(LLVMInitializePTXTargetInfo)
|
_wrap_none2none(LLVMInitializePTXTargetInfo)
|
||||||
_wrap_none2none( LLVMInitializePTXTargetMC )
|
_wrap_none2none( LLVMInitializePTXTargetMC )
|
||||||
_wrap_none2none(LLVMInitializePTXAsmPrinter)
|
_wrap_none2none(LLVMInitializePTXAsmPrinter)
|
||||||
|
#endif
|
||||||
|
|
||||||
/*===----------------------------------------------------------------------===*/
|
/*===----------------------------------------------------------------------===*/
|
||||||
/* Passes */
|
/* Passes */
|
||||||
|
|
@ -1822,11 +1833,17 @@ static PyMethodDef core_methods[] = {
|
||||||
|
|
||||||
_method( LLVMInitializeNativeTarget )
|
_method( LLVMInitializeNativeTarget )
|
||||||
_method( LLVMInitializeNativeTargetAsmPrinter )
|
_method( LLVMInitializeNativeTargetAsmPrinter )
|
||||||
|
#if LLVM_HAS_NVPTX
|
||||||
|
_method( LLVMInitializeNVPTXTarget )
|
||||||
|
_method( LLVMInitializeNVPTXTargetInfo )
|
||||||
|
_method( LLVMInitializeNVPTXTargetMC )
|
||||||
|
_method( LLVMInitializeNVPTXAsmPrinter )
|
||||||
|
#else
|
||||||
_method( LLVMInitializePTXTarget )
|
_method( LLVMInitializePTXTarget )
|
||||||
_method( LLVMInitializePTXTargetInfo )
|
_method( LLVMInitializePTXTargetInfo )
|
||||||
_method( LLVMInitializePTXTargetMC )
|
_method( LLVMInitializePTXTargetMC )
|
||||||
_method( LLVMInitializePTXAsmPrinter )
|
_method( LLVMInitializePTXAsmPrinter )
|
||||||
|
#endif
|
||||||
/* Passes */
|
/* Passes */
|
||||||
|
|
||||||
/*
|
/*
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
21
llvm/core.py
21
llvm/core.py
|
|
@ -2155,8 +2155,19 @@ if _core.LLVMInitializeNativeTargetAsmPrinter():
|
||||||
# should user trigger the initialization?
|
# should user trigger the initialization?
|
||||||
raise llvm.LLVMException("No native asm printer!?")
|
raise llvm.LLVMException("No native asm printer!?")
|
||||||
|
|
||||||
if True: # use PTX
|
|
||||||
_core.LLVMInitializePTXTarget()
|
HAS_PTX = HAS_NVPTX = False
|
||||||
_core.LLVMInitializePTXTargetInfo()
|
if True: # use PTX?
|
||||||
_core.LLVMInitializePTXTargetMC()
|
try:
|
||||||
_core.LLVMInitializePTXAsmPrinter()
|
_core.LLVMInitializePTXTarget()
|
||||||
|
_core.LLVMInitializePTXTargetInfo()
|
||||||
|
_core.LLVMInitializePTXTargetMC()
|
||||||
|
_core.LLVMInitializePTXAsmPrinter()
|
||||||
|
HAS_PTX = True
|
||||||
|
except AttributeError:
|
||||||
|
_core.LLVMInitializeNVPTXTarget()
|
||||||
|
_core.LLVMInitializeNVPTXTargetInfo()
|
||||||
|
_core.LLVMInitializeNVPTXTargetMC()
|
||||||
|
_core.LLVMInitializeNVPTXAsmPrinter()
|
||||||
|
HAS_NVPTX = True
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -893,6 +893,7 @@ LLVMModuleRef LLVMGetModuleFromBitcode(const char *bitcode, unsigned bclen,
|
||||||
return wrap(modulep);
|
return wrap(modulep);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#if LLVM_VERSION_MAJOR <= 3 && LLVM_VERSION_MINOR < 2
|
||||||
unsigned LLVMLinkModules(LLVMModuleRef dest, LLVMModuleRef src, unsigned int mode,
|
unsigned LLVMLinkModules(LLVMModuleRef dest, LLVMModuleRef src, unsigned int mode,
|
||||||
char **out)
|
char **out)
|
||||||
{
|
{
|
||||||
|
|
@ -909,6 +910,7 @@ unsigned LLVMLinkModules(LLVMModuleRef dest, LLVMModuleRef src, unsigned int mod
|
||||||
|
|
||||||
return 1;
|
return 1;
|
||||||
}
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
unsigned char *LLVMGetBitcodeFromModule(LLVMModuleRef module, unsigned *lenp)
|
unsigned char *LLVMGetBitcodeFromModule(LLVMModuleRef module, unsigned *lenp)
|
||||||
{
|
{
|
||||||
|
|
|
||||||
13
llvm/extra.h
13
llvm/extra.h
|
|
@ -37,10 +37,20 @@
|
||||||
#ifndef LLVM_PY_EXTRA_H
|
#ifndef LLVM_PY_EXTRA_H
|
||||||
#define LLVM_PY_EXTRA_H
|
#define LLVM_PY_EXTRA_H
|
||||||
|
|
||||||
|
// select PTX or NVPTX
|
||||||
|
|
||||||
|
#if LLVM_VERSION_MAJOR >= 3 && LLVM_VERSION_MINOR >= 2
|
||||||
|
#define LLVM_HAS_NVPTX 1
|
||||||
|
#else
|
||||||
|
#define LLVM_HAS_NVPTX 0
|
||||||
|
#endif
|
||||||
|
|
||||||
|
|
||||||
#include "llvm-c/Transforms/PassManagerBuilder.h"
|
#include "llvm-c/Transforms/PassManagerBuilder.h"
|
||||||
|
|
||||||
#include "llvm_c_extra.h"
|
#include "llvm_c_extra.h"
|
||||||
|
|
||||||
|
|
||||||
#ifdef __cplusplus
|
#ifdef __cplusplus
|
||||||
extern "C" {
|
extern "C" {
|
||||||
#endif
|
#endif
|
||||||
|
|
@ -381,12 +391,13 @@ LLVMModuleRef LLVMGetModuleFromAssembly(const char *asmtxt, char **out);
|
||||||
LLVMModuleRef LLVMGetModuleFromBitcode(const char *bc, unsigned bclen,
|
LLVMModuleRef LLVMGetModuleFromBitcode(const char *bc, unsigned bclen,
|
||||||
char **out);
|
char **out);
|
||||||
|
|
||||||
|
#if LLVM_VERSION_MAJOR <= 3 && LLVM_VERSION_MINOR < 2
|
||||||
/* Wraps llvm::Linker::LinkModules(). Returns 0 on failure (with errmsg
|
/* Wraps llvm::Linker::LinkModules(). Returns 0 on failure (with errmsg
|
||||||
* filled in) and 1 on success. Dispose error message after use with
|
* filled in) and 1 on success. Dispose error message after use with
|
||||||
* LLVMDisposeMessage(). */
|
* LLVMDisposeMessage(). */
|
||||||
unsigned LLVMLinkModules(LLVMModuleRef dest, LLVMModuleRef src,
|
unsigned LLVMLinkModules(LLVMModuleRef dest, LLVMModuleRef src,
|
||||||
unsigned int, char **errmsg);
|
unsigned int, char **errmsg);
|
||||||
|
#endif
|
||||||
/* Returns pointer to a heap-allocated block of `*len' bytes containing bit code
|
/* Returns pointer to a heap-allocated block of `*len' bytes containing bit code
|
||||||
* for the given module. NULL on error. */
|
* for the given module. NULL on error. */
|
||||||
unsigned char *LLVMGetBitcodeFromModule(LLVMModuleRef module, unsigned *len);
|
unsigned char *LLVMGetBitcodeFromModule(LLVMModuleRef module, unsigned *len);
|
||||||
|
|
|
||||||
37
setup.py
37
setup.py
|
|
@ -29,7 +29,7 @@
|
||||||
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||||
#
|
#
|
||||||
|
|
||||||
import sys, os
|
import sys, os, re
|
||||||
from distutils.core import setup, Extension
|
from distutils.core import setup, Extension
|
||||||
|
|
||||||
LLVM_PY_VERSION = '0.8.2'
|
LLVM_PY_VERSION = '0.8.2'
|
||||||
|
|
@ -72,6 +72,20 @@ def get_llvm_config():
|
||||||
|
|
||||||
return (lc, True)
|
return (lc, True)
|
||||||
|
|
||||||
|
def get_version(llvm_config):
|
||||||
|
# get version number; treat it as fixed point
|
||||||
|
re_version = re.compile(r'(\d+)\.(\d+)')
|
||||||
|
raw = _run(llvm_config + ' --version')
|
||||||
|
major, minor = map(int, re_version.match(raw).groups())
|
||||||
|
return major, minor
|
||||||
|
|
||||||
|
def auto_intrinsic_gen(llvm_config, incdir):
|
||||||
|
# let's do auto intrinsic generation
|
||||||
|
print("Generate intrinsic IDs")
|
||||||
|
from tools import intrgen
|
||||||
|
path = "%s/llvm/Intrinsics.gen" % incdir
|
||||||
|
with open('llvm/_intrinsic_ids.py', 'w') as fout:
|
||||||
|
intrgen.gen(path, fout)
|
||||||
|
|
||||||
def call_setup(llvm_config):
|
def call_setup(llvm_config):
|
||||||
|
|
||||||
|
|
@ -79,7 +93,26 @@ def call_setup(llvm_config):
|
||||||
libdir = _run(llvm_config + ' --libdir')
|
libdir = _run(llvm_config + ' --libdir')
|
||||||
ldflags = _run(llvm_config + ' --ldflags')
|
ldflags = _run(llvm_config + ' --ldflags')
|
||||||
|
|
||||||
ptx_components = ['ptx', 'ptxasmprinter', 'ptxcodegen', 'ptxdesc', 'ptxinfo']
|
llvm_version = get_version(llvm_config)
|
||||||
|
print('LLVM version = %d.%d' % llvm_version)
|
||||||
|
|
||||||
|
auto_intrinsic_gen(llvm_config, incdir)
|
||||||
|
|
||||||
|
if llvm_version <= (3, 1): # select between PTX & NVPTX
|
||||||
|
print('Using PTX')
|
||||||
|
ptx_components = ['ptx',
|
||||||
|
'ptxasmprinter',
|
||||||
|
'ptxcodegen',
|
||||||
|
'ptxdesc',
|
||||||
|
'ptxinfo']
|
||||||
|
else:
|
||||||
|
print('Using NVPTX')
|
||||||
|
ptx_components = ['nvptx',
|
||||||
|
'nvptxasmprinter',
|
||||||
|
'nvptxcodegen',
|
||||||
|
'nvptxdesc',
|
||||||
|
'nvptxinfo']
|
||||||
|
|
||||||
|
|
||||||
libs_core, objs_core = get_libs_and_objs(llvm_config,
|
libs_core, objs_core = get_libs_and_objs(llvm_config,
|
||||||
['core', 'analysis', 'scalaropts', 'executionengine',
|
['core', 'analysis', 'scalaropts', 'executionengine',
|
||||||
|
|
|
||||||
|
|
@ -19,17 +19,22 @@ class TestTargetMachines(unittest.TestCase):
|
||||||
self.assertIn('foo', tm.emit_assembly(m).decode('utf-8'))
|
self.assertIn('foo', tm.emit_assembly(m).decode('utf-8'))
|
||||||
|
|
||||||
def test_ptx(self):
|
def test_ptx(self):
|
||||||
|
if HAS_PTX:
|
||||||
|
arch = 'ptx64'
|
||||||
|
elif HAS_NVPTX:
|
||||||
|
arch = 'nvptx64'
|
||||||
|
else:
|
||||||
|
return # skip this test
|
||||||
m, func = self._build_module()
|
m, func = self._build_module()
|
||||||
func.calling_convention = CC_PTX_KERNEL # set calling conv
|
func.calling_convention = CC_PTX_KERNEL # set calling conv
|
||||||
ptxtm = TargetMachine.lookup(arch='ptx64', cpu='compute_20',
|
ptxtm = TargetMachine.lookup(arch=arch, cpu='compute_20')
|
||||||
features='-double')
|
|
||||||
self.assertTrue(ptxtm.triple)
|
self.assertTrue(ptxtm.triple)
|
||||||
self.assertTrue(ptxtm.cpu)
|
self.assertTrue(ptxtm.cpu)
|
||||||
self.assertTrue(ptxtm.feature_string)
|
|
||||||
ptxasm = ptxtm.emit_assembly(m).decode('utf-8')
|
ptxasm = ptxtm.emit_assembly(m).decode('utf-8')
|
||||||
self.assertIn('foo', ptxasm)
|
self.assertIn('foo', ptxasm)
|
||||||
self.assertIn('map_f64_to_f32', ptxasm)
|
if HAS_NVPTX:
|
||||||
self.assertIn('compute_10', ptxasm)
|
self.assertIn('.address_size 64', ptxasm)
|
||||||
|
self.assertIn('compute_20', ptxasm)
|
||||||
|
|
||||||
def _build_module(self):
|
def _build_module(self):
|
||||||
m = Module.new('TestTargetMachines')
|
m = Module.new('TestTargetMachines')
|
||||||
|
|
|
||||||
0
tools/__init__.py
Normal file
0
tools/__init__.py
Normal file
|
|
@ -6,7 +6,7 @@
|
||||||
|
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
def gen(f):
|
def gen(f, out=sys.stdout):
|
||||||
intr = []
|
intr = []
|
||||||
maxw = 0
|
maxw = 0
|
||||||
flag = False
|
flag = False
|
||||||
|
|
@ -26,7 +26,8 @@ def gen(f):
|
||||||
idx = 1
|
idx = 1
|
||||||
for i in intr:
|
for i in intr:
|
||||||
s = 'INTR_' + i.upper()
|
s = 'INTR_' + i.upper()
|
||||||
print('%s = %d' % (s.ljust(maxw), idx))
|
out.write('%s = %d\n' % (s.ljust(maxw), idx))
|
||||||
idx += 1
|
idx += 1
|
||||||
|
|
||||||
gen(sys.argv[1])
|
if __name__ == '__main__':
|
||||||
|
gen(sys.argv[1])
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue