使用numba jit装饰器编译函数时调用numpy.sort报错如何解决
问题原因
你遇到的报错核心有两点:
- 数组初始化语法错误:你写的
dtype = np.dtype((int,int))会创建每个元素为双int的结构化数组,而非普通的二维int32数组,如需创建int32类型的二维数组,直接改为dtype=np.int32即可。 - Numba的nopython模式下,当前仅支持对一维数组调用
np.sort,你代码中传入的是二维数组,没有对应的重载实现,因此触发编译失败。你看到官方文档标注np.sort受支持是准确的,只是未覆盖二维入参的场景。
解决方案
根据不同的排序需求可选择对应实现:
- 需求为排序二维数组的所有元素:先将数组展平为一维再调用
np.sort,示例代码如下:
import numba as nb import numpy as np gg = np.array([[1,0,2],[1,2,1]], dtype = np.int32) @nb.njit(nb.void(nb.int32[:,:])) def kk (gg): # 展平为一维后排序,如有需要可reshape回原形状 sorted_res = np.sort(gg.flatten()) # 若需要修改原数组,可补充回写逻辑:gg[:] = sorted_res.reshape(gg.shape)
- 需求为按行/按列分别排序:遍历对应维度,逐行/逐列传入一维数组调用
np.sort,以按行排序为例:
@nb.njit(nb.void(nb.int32[:,:])) def kk (gg): for row_idx in range(gg.shape[0]): gg[row_idx] = np.sort(gg[row_idx])
其他说明
如果强制使用object模式运行,可以在装饰器中添加forceobj=True参数,但此时numba无法对代码做高性能编译优化,失去加速意义,不推荐使用。
内容的提问来源于stack exchange,提问作者user2625363
相关产品推荐
相关产品推荐

