如何在Numpy中获取递增序列中浮点数的区间索引?
使用numpy.searchsorted()解决区间索引匹配问题
嘿,你找的这个功能numpy里真的有现成实现!就是numpy.searchsorted()——专门解决这种在已排序数组里找区间索引的问题,完全符合你的需求。
核心原理
因为你的t数组是已排序且互不相同的,searchsorted()会为T中的每个元素,找到它应该插入t中的位置(保持t排序的位置)。对于左闭右开区间[t_i, t_{i+1})来说:
- 插入位置减1,就是我们要找的
i_j——因为插入位置是第一个大于T_j的t元素的索引,那T_j必然落在前一个元素到这个元素的左闭右开区间里。
代码示例
import numpy as np # 示例输入 t = np.array([1, 3, 5, 7, 9]) # 已排序的p+1个元素(p=4) T = np.array([2, 3, 5, 8, 9, 1]) # 所有元素属于[t0, tp] # 计算索引数组 indices = np.searchsorted(t, T) - 1 print(indices) # 输出:[0 1 2 3 4 0]
边界情况处理
注意上面的结果里,T中的9(等于t[-1])得到的索引是4,但t[4]是9,不存在t[5],所以[t_4, t_5)是无效区间。如果你的需求是把等于t_p的元素归到最后一个有效区间[t_{p-1}, t_p)(或者说视为左闭右闭的[t_{p-1}, t_p]),可以额外加一步处理:
# 修正等于tₚ的元素的索引 indices[T == t[-1]] = len(t) - 2 print(indices) # 输出:[0 1 2 3 3 0]
额外说明
- 不需要提前排序
T:searchsorted()会逐个处理T中的每个元素,不管T是否排序。当然如果T是排序的,函数内部可以用更高效的算法,但不是必须的。 - 性能优势:这是numpy的内置矢量化操作,比手动写循环快几个数量级,尤其适合大规模的
T数组。
内容的提问来源于stack exchange,提问作者bela83
相关产品推荐
相关产品推荐

