如何为接收np.array输入的Numba函数正确设置签名
Numba函数签名设置错误的解决方法
问题场景
无签名版本的Numba函数运行正常:
import numpy as np import numba as nb import time mat = np.random.normal(0, 1, size=(1000000, 10)) @nb.jit(nopython=True) def f(mat): max_min = 0 for i in range(mat.shape[0]): max_min += mat[i].max() - mat[i].min() return max_min / mat.shape[0] start = time.time() print(f(mat)) end = time.time() print(end - start)
但设置签名为@nb.jit(nb.float64(nb.float64), nopython=True)后,出现如下报错:
TypingError: Failed in nopython mode pipeline (step: nopython frontend) Unknown attribute 'shape' of type float64
错误原因
你设置的签名nb.float64(nb.float64)表示函数接受单个float64类型的标量,返回float64标量,但实际传入的是二维numpy数组。Numba按照签名将mat解析为float64标量,而标量没有shape属性,因此触发报错。
解决方法
需要将签名修改为匹配二维float64数组的类型,使用Numba的数组类型语法nb.float64[:, :]表示二维数组:
import numpy as np import numba as nb import time mat = np.random.normal(0, 1, size=(1000000, 10)) @nb.jit(nb.float64(nb.float64[:, :]), nopython=True) def f(mat): max_min = 0 for i in range(mat.shape[0]): max_min += mat[i].max() - mat[i].min() return max_min / mat.shape[0] start = time.time() print(f(mat)) end = time.time() print(end - start)
如果需要兼容不同维度的数组,可以使用灵活的数组类型nb.float64[:](任意维度),但明确指定维度能让Numba做更精准的优化。
内容的提问来源于stack exchange,提问作者Zhaoyan Zhou
相关产品推荐
相关产品推荐

