如何用Numba njit获取含至少一个NaN的NumPy数组行索引?
解决Numba njit中检测含NaN行索引的问题
由于Numba暂不支持带axis参数的np.any()调用,你可以通过以下两种兼容Numba njit的方式实现需求:
方法一:手动遍历行列检查
通过双层循环逐行逐元素检测NaN,找到目标行后立即终止该行的检查以提升效率:
import numba as nb import numpy as np @nb.njit def get_nan_row_indices(X): row_count, col_count = X.shape has_nan = np.zeros(row_count, dtype=np.bool_) for i in range(row_count): for j in range(col_count): if np.isnan(X[i, j]): has_nan[i] = True break # 找到NaN后停止检查当前行 return has_nan # 测试 X = np.array([[1, 2], [3, np.nan], [6, 9]]) print(get_nan_row_indices(X)) # 输出:array([False, True, False])
方法二:利用一维数组的np.any
Numba支持对一维数组调用np.any(),因此可以遍历每行并单独检查该行是否包含NaN:
@nb.njit def get_nan_row_indices_v2(X): row_count = X.shape[0] has_nan = np.zeros(row_count, dtype=np.bool_) for i in range(row_count): if np.any(np.isnan(X[i])): has_nan[i] = True return has_nan # 测试 X = np.array([[1, 2], [3, np.nan], [6, 9]]) print(get_nan_row_indices_v2(X)) # 输出:array([False, True, False])
两种方法都能得到与纯NumPy实现一致的结果,其中方法一在NaN出现位置靠前的场景下效率更高,方法二代码更简洁易读。
内容的提问来源于stack exchange,提问作者Keptain
相关产品推荐
相关产品推荐

