如何基于一维断点集快速确定数值所属区间?
用Numpy快速定位一维值所属区间
对于任意数量的一维断点,Numpy的digitize()函数可以完美解决区间定位问题,无需嵌套if语句,代码简洁且效率高。
核心用法说明
numpy.digitize(x, bins, right=False)会返回每个x值在升序断点数组bins中的插入位置,这个位置直接对应所属区间的索引,再通过索引映射到自定义的区间标识即可。
- 默认
right=False时,区间规则为:- 索引0 →
x < bins[0] - 索引i(1≤i<len(bins)) →
bins[i-1] ≤ x < bins[i] - 索引len(bins) →
x ≥ bins[-1]
- 索引0 →
- 若需左开右闭区间,设置
right=True即可,对应规则会调整为bins[i-1] < x ≤ bins[i](首尾区间同理)。
代码示例
import numpy as np # 定义断点数组(确保升序,若无序可先用np.sort()排序) breakpoints = np.array([-1, 0, 1]) # 自定义区间标识,数量为断点数+1 interval_labels = ['A', 'B', 'C', 'D'] # 对应关系:A→x<-1, B→-1≤x<0, C→0≤x<1, D→x≥1 # 测试批量x值 test_x = np.array([-2, -0.5, 0.3, 2, -1, 1]) # 获取每个x对应的区间索引 interval_indices = np.digitize(test_x, breakpoints) # 映射为区间标识 result = [interval_labels[idx] for idx in interval_indices] print(result) # 输出: ['A', 'B', 'C', 'D', 'B', 'D']
注意事项
- 断点数组必须是升序排列,如果原始断点无序,先执行
breakpoints = np.sort(breakpoints)处理。 - 单个x值也适用,无需批量输入,直接传入单个数值即可得到对应索引。
内容的提问来源于stack exchange,提问作者irene
相关产品推荐
相关产品推荐

