Numba的nopython模式下numpy数组排序编译失败是什么原因?
问题原因及解决方案
你遇到的编译报错本质是当前使用的Numba版本对高维数组排序操作的支持存在限制,具体解决方式如下:
核心原因
Numba 0.57.0之前的版本仅支持一维数组的np.sort、np.argsort、数组原生方法.sort()、.argsort()调用,传入二维及以上维度的数组时会出现函数签名匹配失败的TypingError,就是你收到的报错内容。
解决方法
方案1:升级Numba版本(推荐)
首先查看当前Numba版本:
import numba print(numba.__version__)
如果版本号低于0.57.0,直接升级到最新稳定版即可原生支持二维数组排序操作:
pip install --upgrade numba
升级完成后你原本的代码无需修改即可正常编译运行。
方案2:低版本兼容写法
如果受环境限制无法升级Numba,可以手动遍历二维数组的每个维度单独做排序操作,替代原本的高维数组直接排序调用,示例如下:
@njit() def accuracy(x,y,z): # 等价于x.argsort(axis=1)的兼容实现 for i in range(x.shape[0]): row_sort_idx = np.argsort(x[i]) x[i] = x[i][row_sort_idx] # 后续原有逻辑保持不变 accuracy_y = int(((np.equal(y, x).mean())*100)%100) accuracy_z = int(((np.equal(z, x).mean())*100)%100) return accuracy_y,accuracy_z
额外注意点
- 数组原生的
.sort()是原地修改操作,会直接改变输入数组的原始值,如果不需要修改原数组,请使用np.sort或者.argsort()生成新数组 - 即使是0.57.0及以上的高版本Numba,目前也不支持排序操作传入
axis=None的拉平排序写法,需要显式指定axis参数,或者手动将数组拉平为一维后再调用排序方法
内容的提问来源于stack exchange,提问作者tjaqu787
相关产品推荐
相关产品推荐

