Numba nopython模式下函数参数报错,如何进行类型标注?
Numba nopython模式下传递函数参数的类型识别问题
问题重现
使用@numba.jit(nopython=True)装饰器编译接收函数作为参数的代码时,会触发TypingError,无法识别函数类型。示例代码如下:
import numba import numpy as np x = np.random.randn(10,10) f = lambda x : (x>0)*x @numba.jit(nopython=True) def a(x,f): return f(x)**2+x a(x,f)
报错信息
Traceback (most recent call last): File "<stdin>", line 1, in <module> File "C:\ProgramData\Anaconda3\envs\pytorch\lib\site-packages\numba\core\dispatcher.py", line 468, in _compile_for_args error_rewrite(e, 'typing') File "C:\ProgramData\Anaconda3\envs\pytorch\lib\site-packages\numba\core\dispatcher.py", line 409, in error_rewrite raise e.with_traceback(None) numba.core.errors.TypingError: Failed in nopython mode pipeline (step: nopython frontend) non-precise type pyobject During: typing of argument at <stdin> (2) File "<stdin>", line 2: <source missing, REPL/exec in use?> This error may have been caused by the following argument(s): - argument 1: Cannot determine Numba type of <class 'function'>
移除nopython=True后代码可正常运行:
@numba.jit def a(x,f): return f(x)**2+x a(x,f)
解决方案
可以通过以下两种方式让Numba在nopython模式下识别函数类型:
1. 提前用@numba.jit编译函数参数f
将作为参数传递的函数f先用@numba.jit编译,Numba就能识别其类型。注意lambda函数无法直接用@numba.jit编译,需要改成普通函数形式:
import numba import numpy as np x = np.random.randn(10,10) # 先编译f @numba.jit def f(x): return (x>0)*x @numba.jit(nopython=True) def a(x,f): return f(x)**2 + x a(x,f)
2. 显式指定函数类型签名
通过给a函数添加类型签名,明确声明f的函数类型,适合输入输出类型固定的场景:
import numba from numba import types import numpy as np # 定义f的类型:输入为float64二维数组,输出同类型 f_type = types.FunctionType(types.float64[:, :](types.float64[:, :])) x = np.random.randn(10,10) @numba.jit def f(x): return (x>0)*x # 给a函数指定类型签名 @numba.jit("float64[:, :](float64[:, :], FunctionType)", nopython=True) def a(x,f): return f(x)**2 + x a(x,f)
原理说明
Numba的nopython模式要求所有变量必须是Numba可识别的静态类型,未编译的Python函数属于pyobject类型,不符合nopython模式的类型要求。提前编译函数参数或显式指定类型,能让Numba在编译阶段确定函数的类型信息,从而通过类型检查。
内容的提问来源于stack exchange,提问作者Uri Cohen
相关产品推荐
相关产品推荐

