如何将含Numba编译函数的数组作为参数传入njit函数?
解决Numba中传入编译函数数组的问题
你遇到的non-precise type array(pyobject, 1d, C)错误,根源是用普通numpy数组存储Numba编译函数时,数组会被自动设为object dtype——而Numba的nopython模式无法处理类型模糊的object数组,因为它需要明确的静态类型才能完成编译优化。
可行解决方案:使用Numba typed.List
Numba提供了typed.List这个类型化容器,专门用来存储能被Numba识别的同类型元素,包括签名一致的编译函数。修改后的代码如下:
import numpy as np from numba import njit from numba.typed import List @njit() def function1(x, y): return x > y @njit() def function2(x, y): return x < y @njit() def main(input_list): print(input_list[0](1, 2)) print(input_list[1](1, 2)) # 创建Numba类型化列表并添加编译后的函数 function_list = List() function_list.append(function1) function_list.append(function2) main(function_list)
这段代码能正常运行,因为typed.List会自动推断存储元素的类型(这里是(int64, int64)->bool的编译函数),Numba可以在nopython模式下直接处理这个容器。
备选方案:固定数量函数直接传参
如果你的函数数量是固定的,直接把每个函数作为单独参数传入njit函数,也是更简单的方式:
@njit() def main(func1, func2): print(func1(1, 2)) print(func2(1, 2)) main(function1, function2)
注意事项
- 不要尝试用
np.array([function1, function2], dtype=function)这种写法,numpy没有原生的函数类型,最终还是会生成object数组,无法被Numba识别。 - 存储在
typed.List中的函数必须有一致的签名(比如都是接受两个int参数返回bool),否则Numba会无法推断统一类型。
内容的提问来源于stack exchange,提问作者CK Chua
相关产品推荐
相关产品推荐

