如何用Numpy正确查找一维数组中近似相等值的子数组?
寻找一维NumPy数组中近似相等值的子数组
我有一个一维NumPy数组,想要找到其中包含近似相等值的子数组。这里的近似相等指子数组中任意两个元素的差值不超过设定的容差(tolerance)。比如数组[1.0, 2.2, 1.4, 1.8, 1.5, 2.1],容差设为0.2,期望结果是[[1.4, 1.5], [2.1, 2.2]];输入数组[1.0, 1.1, 1.2, 1.3, 1.4]、容差0.2时,期望输出应为[[1.0, 1.1, 1.2], [1.1, 1.2, 1.3], [1.2, 1.3, 1.4]]。
原函数的问题
你编写的函数存在两个主要问题:
- 最后一步的子集删除逻辑错误:循环中删除列表元素会导致索引混乱,错误地移除了需要保留的重叠子数组。
- 逻辑不符合实际需求:原函数生成的子数组是基于每个元素作为比较基准,无法得到你期望的最长连续近似相等子数组。
优化后的NumPy实现
利用滑动窗口算法可以高效解决这个问题,核心思路是在排序后的数组中,找到所有满足max - min ≤ tol的最长连续子数组:
import numpy as np def find_max_consecutive_almost_equal(arr, tol): sorted_arr = np.sort(arr) n = len(sorted_arr) result = [] max_len = 0 left = 0 for right in range(n): # 当窗口内最大最小差值超过容差时,移动左边界 while sorted_arr[right] - sorted_arr[left] > tol: left += 1 current_window_len = right - left + 1 # 更新最长子数组列表 if current_window_len >= 2: if current_window_len > max_len: max_len = current_window_len result = [sorted_arr[left:right+1].tolist()] elif current_window_len == max_len: sub_arr = sorted_arr[left:right+1].tolist() if sub_arr not in result: result.append(sub_arr) return result
测试验证
- 测试案例1:
test1 = np.array([1.0, 2.2, 1.4, 1.8, 1.5, 2.1]) tolerance = 0.2 print(find_max_consecutive_almost_equal(test1, tolerance)) # 输出: [[1.4, 1.5], [2.1, 2.2]] - 测试案例2:
test2 = np.array([1.0, 1.1, 1.2, 1.3, 1.4]) tolerance = 0.2 print(find_max_consecutive_almost_equal(test2, tolerance)) # 输出: [[1.0, 1.1, 1.2], [1.1, 1.2, 1.3], [1.2, 1.3, 1.4]] - 测试案例3:
test3 = np.array([2.6, 1.2, 1.5, 1.8, 2.0, 2.2, 2.5, 1.1, 1.4]) tolerance = 0.15 print(find_max_consecutive_almost_equal(test3, tolerance)) # 输出: [[1.1, 1.2], [1.4, 1.5], [2.5, 2.6]]
另一种实现:满足"存在中心点"的定义
如果严格按照你最初的定义(子数组中所有值与某个中心点的差值≤tol,等价于max - min ≤ 2*tol),可以用二分查找快速定位每个元素的有效范围,再去重:
import numpy as np def find_almost_equal_with_center(arr, tol): sorted_arr = np.sort(arr) n = len(sorted_arr) result = [] for x in sorted_arr: # 二分查找左右边界 left = np.searchsorted(sorted_arr, x - tol, side='left') right = np.searchsorted(sorted_arr, x + tol, side='right') - 1 sub_arr = sorted_arr[left:right+1] if len(sub_arr) >= 2: result.append(sub_arr.tolist()) # 去重并排序 unique_result = list(set(tuple(r) for r in result)) unique_result = sorted([list(r) for r in unique_result], key=lambda x: x[0]) return unique_result
内容的提问来源于stack exchange,提问作者Frank Tap
相关产品推荐
相关产品推荐

