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

如何让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列场景,正确流程如下:

  1. 对每一列排序并记录排序后的原始行索引
  2. 对每组(start, stop),在每一列的排序结果中找到匹配的索引范围
  3. 取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 15:58:13