94 lines
2.2 KiB
Python
94 lines
2.2 KiB
Python
# coding: utf-8
|
|
# pylint: disable=invalid-name, no-member
|
|
""" ctypes library of nnvm and helper functions """
|
|
from __future__ import absolute_import
|
|
|
|
import sys
|
|
import ctypes
|
|
import numpy as np
|
|
from . import libinfo
|
|
|
|
__all__ = ['TVMError']
|
|
#----------------------------
|
|
# library loading
|
|
#----------------------------
|
|
if sys.version_info[0] == 3:
|
|
string_types = str,
|
|
numeric_types = (float, int, np.float32, np.int32)
|
|
# this function is needed for python3
|
|
# to convert ctypes.char_p .value back to python str
|
|
py_str = lambda x: x.decode('utf-8')
|
|
else:
|
|
string_types = basestring,
|
|
numeric_types = (float, int, long, np.float32, np.int32)
|
|
py_str = lambda x: x
|
|
|
|
|
|
class TVMError(Exception):
|
|
"""Error that will be throwed by all functions"""
|
|
pass
|
|
|
|
def _load_lib():
|
|
"""Load libary by searching possible path."""
|
|
lib_path = libinfo.find_lib_path()
|
|
lib = ctypes.CDLL(lib_path[0], ctypes.RTLD_GLOBAL)
|
|
# DMatrix functions
|
|
lib.TVMGetLastError.restype = ctypes.c_char_p
|
|
return lib
|
|
|
|
# version number
|
|
__version__ = libinfo.__version__
|
|
# library instance of nnvm
|
|
_LIB = _load_lib()
|
|
|
|
#----------------------------
|
|
# helper function definition
|
|
#----------------------------
|
|
def check_call(ret):
|
|
"""Check the return value of C API call
|
|
|
|
This function will raise exception when error occurs.
|
|
Wrap every API call with this function
|
|
|
|
Parameters
|
|
----------
|
|
ret : int
|
|
return value from API calls
|
|
"""
|
|
if ret != 0:
|
|
raise TVMError(py_str(_LIB.TVMGetLastError()))
|
|
|
|
|
|
def c_str(string):
|
|
"""Create ctypes char * from a python string
|
|
Parameters
|
|
----------
|
|
string : string type
|
|
python string
|
|
|
|
Returns
|
|
-------
|
|
str : c_char_p
|
|
A char pointer that can be passed to C API
|
|
"""
|
|
return ctypes.c_char_p(string.encode('utf-8'))
|
|
|
|
|
|
def c_array(ctype, values):
|
|
"""Create ctypes array from a python array
|
|
|
|
Parameters
|
|
----------
|
|
ctype : ctypes data type
|
|
data type of the array we want to convert to
|
|
|
|
values : tuple or list
|
|
data content
|
|
|
|
Returns
|
|
-------
|
|
out : ctypes array
|
|
Created ctypes array
|
|
"""
|
|
return (ctype * len(values))(*values)
|