Numba njit模式下sorted排序含np.inf数组结果与Python不一致原因
问题背景
现有一个形状为n×3的numpy ndarray数组,需求为按照数组第三列的值进行降序排序,该列中可能包含np.inf值。最初采用Python内置sorted函数实现排序逻辑,但在为核心函数添加@njit装饰器后,得到的排序结果与原生Python环境下的运行结果存在明显差异。
用于复现问题的测试代码如下:
#Sorting Functions import numpy as np from numba import njit @njit def sort_me_numba(arr): res = sorted(arr, key=lambda x: x[2], reverse=True) return res def sort_me_python(arr): res = sorted(arr, key=lambda x: x[2], reverse=True) return res #Making data format I have arr = np.concatenate([np.random.normal(loc=0.1, scale=0.005, size=1_000).reshape((-1, 1)) for i in range(3)], axis=1) samples = [0, 14, 53, 344, 43, 654, 435, 33] arr[samples, 2] = np.inf print('-----------NUMBA with numpy array------------') sort_me_numba(arr) print('-----------PYTHON with numpy array------------') sort_me_python(arr) print('-----------PYTHON with list------------') sort_me_python(list(arr))
运行上述测试代码可复现结果不一致的现象,但当测试代码中samples = [0, 1, 2, 442]时,两种运行模式下的排序结果又完全一致。
现象产生原因
- Numba在
@njit模式下迭代numpy二维数组时,不会为每一行生成独立的数组对象,而是复用同一个临时内存视图存储当前迭代到的行,迭代过程中该视图的内容会被不断覆写为下一行的数据。 - 原生Python实现的
sorted在遍历元素时,会提前执行key函数、缓存每个元素对应的key值,后续排序比较直接使用缓存的key,不会再回头访问原元素。但Numba编译版本的sorted没有做key缓存,排序比较阶段才会临时去取元素的x[2]值,此时临时视图早已被覆写为迭代过程中最后访问的行内容,拿到的根本不是对应行原本的第三列值,自然会出现排序结果错乱。 - 当
samples取值为[0, 1, 2, 442]时结果一致完全是巧合:迭代过程中临时视图多次被覆写为第三列为np.inf的行,排序阶段取key时恰好拿到了np.inf值,和预期排序规则的结果撞车,并不是排序逻辑正确。
正确实现方式
在Numba中不要直接套用Python的sorted加lambda key的写法处理numpy数组行排序,直接使用numpy原生排序接口即可,Numba可以对这类接口做完美编译加速:
@njit def sort_me_numba_correct(arr): # 取第三列降序排序的索引,重排数组 sort_idx = np.argsort(-arr[:, 2]) return arr[sort_idx]
内容的提问来源于stack exchange,提问作者Eugene
相关产品推荐
相关产品推荐

