如何将Numba即时编译函数列表传入另一Numba即时编译函数?
问题描述
以下普通Python代码运行正常:
import numpy as np import numba as nb def f1(x): return 2 * x def f2(x): return x - 4 def f(funcs, x): out = np.zeros(len(funcs)) for i in range(len(out)): out[i] = funcs[i](x) return out f([f1, f2], 3) >>> array([ 6., -1.])
但为每个函数添加@nb.njit装饰器后,运行会报错:
TypeError: can't unbox heterogeneous list: type(CPUDispatcher(<function f1 at 0x2a10455e0>)) != type(CPUDispatcher(<function f2 at 0x2a5dedc10>))
原因是Numba无法识别异构函数列表的类型——每个被njit编译后的函数是独立的CPUDispatcher实例,类型被视为不同。
需要解决的问题:如何将即时编译后的函数列表传入另一即时编译函数,让Numba能正常识别、编译并运行?
添加装饰器后无法运行的代码:
@nb.njit def f1(x): return 2 * x @nb.njit def f2(x): return x - 4 @nb.njit def f(funcs, x): out = np.zeros(len(funcs)) for i in range(len(out)): out[i] = funcs[i](x) return out
解决方案
方法1:使用Numba类型化列表统一函数类型
Numba的typed.List要求列表内元素类型一致,先定义函数的类型签名,再把编译后的函数添加到类型化列表中:
import numpy as np import numba as nb from numba import typed # 定义函数类型:输入float,返回float func_type = nb.types.FunctionType(nb.float64(nb.float64)) @nb.njit(nb.float64(nb.float64)) def f1(x): return 2 * x @nb.njit(nb.float64(nb.float64)) def f2(x): return x - 4 @nb.njit def f(funcs, x): out = np.zeros(len(funcs)) for i in range(len(out)): out[i] = funcs[i](x) return out # 创建类型化列表并添加函数 funcs_list = typed.List.empty_list(func_type) funcs_list.append(f1) funcs_list.append(f2) # 调用测试 print(f(funcs_list, 3.0)) # 输出: [ 6. -1.]
要点:给每个编译函数显式指定签名,确保类型一致;用typed.List存储函数,让Numba识别列表的统一类型。
方法2:将函数封装为Numba jitable类
把函数逻辑封装到类的方法中,利用类的统一性让Numba识别类型:
import numpy as np import numba as nb @nb.experimental.jitclass class FuncWrapper: def __init__(self, func_type): self.func_type = func_type # 用标识区分不同函数逻辑 def apply(self, x): if self.func_type == 1: return 2 * x elif self.func_type == 2: return x - 4 @nb.njit def f(funcs, x): out = np.zeros(len(funcs)) for i in range(len(out)): out[i] = funcs[i].apply(x) return out # 创建封装实例 f1_wrap = FuncWrapper(1) f2_wrap = FuncWrapper(2) # 调用测试 print(f([f1_wrap, f2_wrap], 3.0)) # 输出: [ 6. -1.]
要点:使用jitclass封装不同逻辑,统一实例类型;通过类方法调用具体逻辑,避免直接传递异构函数对象。
方法3:使用Numba的generated_jit处理多函数情况
如果函数逻辑差异较大,用generated_jit根据输入函数生成对应的编译代码:
import numpy as np import numba as nb from numba import typed @nb.njit def f1(x): return 2 * x @nb.njit def f2(x): return x - 4 @nb.generated_jit def f(funcs, x): # 针对传入的函数列表生成特定代码 def impl(funcs, x): out = np.zeros(len(funcs)) for i in range(len(out)): out[i] = funcs[i](x) return out return impl # 创建类型化列表存储函数 funcs_list = typed.List() funcs_list.append(f1) funcs_list.append(f2) # 调用测试 print(f(funcs_list, 3.0)) # 输出: [ 6. -1.]
要点:generated_jit允许根据输入类型动态生成编译逻辑;仍需配合typed.List确保函数列表类型统一。
内容的提问来源于stack exchange,提问作者PyRsquared
相关产品推荐
相关产品推荐

