如何为接收结构化数组的Numba函数指定签名?
为接收结构化数组的Numba函数指定签名的正确方法
要给接收结构化数组的Numba函数指定显式签名,核心是用Numba原生类型定义数组参数,而非直接使用numpy的dtype。以下是两种可行写法:
方法1:用nb.types.Array构造数组类型
通过nb.types.Array明确指定数组的元素类型、维度和内存布局:
import numba as nb import numpy as np PairSpec = [("x", np.float32), ("y", np.float32)] Pair = np.dtype(PairSpec) # 将numpy结构化dtype转换为Numba原生类型 NumbaPair = nb.from_dtype(Pair) # 签名格式:返回类型(参数类型),用nb.types.Array定义结构化数组类型 @nb.jit(nb.float32(nb.types.Array(NumbaPair, 1, 'C'))) def sum_pairs(pairs): pair = pairs[0] return pair.x + pair.y pairs = np.array([(2, 3)], dtype=PairSpec) print(sum_pairs(pairs)) # 输出5.0
nb.types.Array(元素类型, 维度, 内存布局):1代表一维数组,'C'表示C连续内存(也可填'A'适配任意布局)。
方法2:用Numba类型的索引语法(简洁写法)
直接通过Numba原生类型加[:]表示一维数组,写法更简洁:
import numba as nb import numpy as np PairSpec = [("x", np.float32), ("y", np.float32)] Pair = np.dtype(PairSpec) NumbaPair = nb.from_dtype(Pair) # 直接用Numba类型+[:]表示一维结构化数组 @nb.jit(nb.float32(NumbaPair[:])) def sum_pairs(pairs): pair = pairs[0] return pair.x + pair.y pairs = np.array([(2, 3)], dtype=PairSpec) print(sum_pairs(pairs)) # 输出5.0
注意事项
- 禁止直接用numpy的
Pair[:]作为签名参数:Numba无法识别numpy dtype的数组语法,必须先用nb.from_dtype转换为Numba原生类型。 - 避免使用Python内置函数名(如
sum)作为自定义函数名,防止命名冲突。
内容的提问来源于stack exchange,提问作者Valéry
相关产品推荐
相关产品推荐

