如何过滤多维数组中含NaN的子数组并获取其原索引?
过滤多维数组中含NaN的子数组并获取原索引
问题分析
你之前用循环删除的方法失败,核心原因是每次调用np.delete后数组长度会缩短,后续循环的索引和原数组的索引不再对应,最终导致索引越界。numpy的优势在于向量化操作,这类问题完全不需要用循环处理。
解决方法
1. 过滤含NaN的子数组
利用numpy的向量化判断生成布尔掩码,再通过布尔索引直接筛选出不含NaN的子数组:
import numpy as np arr = np.array([[1,2,3], [np.nan, np.nan, np.nan], [4,5,6]]) # 生成布尔数组:True表示对应子数组包含NaN,False表示不包含 has_nan = np.isnan(arr).any(axis=1) # 取反筛选出不含NaN的子数组 filtered_arr = arr[~has_nan] print(filtered_arr) # 输出:[[1 2 3] # [4 5 6]]
2. 获取被过滤子数组的原索引
如果需要知道原数组中哪些子数组被过滤,用np.where提取对应索引即可:
deleted_indices = np.where(has_nan)[0] print(deleted_indices) # 输出:[1]
循环删除出错的原因
你的循环逻辑中,第一次删除索引1的元素后,原数组长度从3变为2,但循环依然会执行到i=2,此时数组的最大有效索引是1,直接触发索引越界错误。numpy的向量化操作一次性完成判断和筛选,完全规避了这类问题。
内容的提问来源于stack exchange,提问作者E. M.
相关产品推荐
相关产品推荐

