You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.30 06:27:13