如何让numpy.searchsorted无合适索引时返回nan而非0或N?
问题描述
numpy.searchsorted() 在无匹配元素时会返回0或数组长度N,导致后续用np.arange()生成索引时得到包含所有元素的子集,不符合需求。
我的场景是基于未排序的pandas DataFrame的3列(对应3个numpy.ndarray),通过二分搜索匹配每一组(start, stop)条件,最终返回原始DataFrame中满足条件的行(需保留其他列数据,因此必须通过索引获取)。当前示例代码如下:
import numpy as np a = np.array([1,2,3,4,5]) start = np.array([0.0,1.0,3.0,3.0,5.0,6.0]) stop = np.array([1.0,3.0,3.0,5.0,6.0,7.0]) starts = np.searchsorted(a, start, side="left") stops = np.searchsorted(a, stop, side="right") length = range(len(starts)) indices = [np.arange(starts[i], stops[i] + 1) for i in length]
我的思路是利用np.argsort()的返回值,对3列的搜索结果取索引交集再子集化,但searchsorted返回0或N的行为会干扰结果。希望实现:当无合适索引时,后续子集化得到空结果,比如让searchsorted返回nan,或有其他可行方法。
解决方案
方法1:修正searchsorted结果,生成空索引
在得到starts和stops后,手动判断无匹配的情况,将对应索引范围设为空数组:
import numpy as np a = np.array([1,2,3,4,5]) start = np.array([0.0,1.0,3.0,3.0,5.0,6.0]) stop = np.array([1.0,3.0,3.0,5.0,6.0,7.0]) starts = np.searchsorted(a, start, side="left") stops = np.searchsorted(a, stop, side="right") indices = [] for s_start, s_stop in zip(starts, stops): # 当起始索引 >= 结束索引时,说明无匹配元素 if s_start >= s_stop: indices.append(np.array([])) else: # 注意:原代码中stops[i]+1会导致取到超出范围的元素,因为searchsorted(side="right")返回的是开区间右边界 indices.append(np.arange(s_start, s_stop)) # 输出结果 for idx in indices: print(idx)
输出:
[] [0 1] [] [2 3] [4] []
该结果符合预期:无匹配的条件(如start=0.0, stop=1.0或start=6.0, stop=7.0)会返回空数组。
方法2:结合排序索引,直接筛选原始行(适配多列场景)
针对你实际的3列场景,正确流程如下:
- 对每一列排序并记录排序后的原始行索引
- 对每组
(start, stop),在每一列的排序结果中找到匹配的索引范围 - 取3列匹配索引的交集,得到原始DataFrame中满足所有条件的行
示例代码:
import numpy as np import pandas as pd # 模拟未排序的DataFrame df = pd.DataFrame({ 'col1': [3,1,4,2,5], 'col2': [7,5,9,6,8], 'col3': [12,10,14,11,13], 'other_col': ['a','b','c','d','e'] }) # 保存每一列的排序索引和排序后的值 sorted_indices = {col: np.argsort(df[col].values) for col in ['col1','col2','col3']} sorted_cols = {col: df[col].values[idx] for col, idx in sorted_indices.items()} # 定义多组条件,每组对应col1、col2、col3的(start, stop) start_conditions = np.array([[0,6,10], [1,5,10], [3,8,12]]) stop_conditions = np.array([[1,8,11], [3,7,13], [5,9,14]]) # 处理每组条件 for start, stop in zip(start_conditions, stop_conditions): matched_sets = [] valid = True for idx, col in enumerate(sorted_cols.keys()): s_start = np.searchsorted(sorted_cols[col], start[idx], side='left') s_stop = np.searchsorted(sorted_cols[col], stop[idx], side='right') if s_start >= s_stop: valid = False break # 取出该列匹配的原始行索引,转成集合方便求交集 matched_idx = sorted_indices[col][s_start:s_stop] matched_sets.append(set(matched_idx)) if not valid: print("无满足条件的行") continue # 取三列的索引交集 common_idx = matched_sets[0] & matched_sets[1] & matched_sets[2] if common_idx: print("满足条件的行:") print(df.loc[list(common_idx)]) else: print("无满足条件的行")
该方法直接处理原始行索引的交集,既避免了无效索引范围的问题,又能正确保留其他列的数据。
内容的提问来源于stack exchange,提问作者Buzz B
相关产品推荐
相关产品推荐

