优化k维数组相似性检测函数,解决大数据集运行效率问题
高效筛选无匹配邻接数组的方案
问题描述
给定一个由长度为k的数组组成的列表,需筛选出所有**不存在其他数组与之满足“每个对应索引元素的差值均在500以内”**的数组。
示例:
输入:
intervals = [[9000, 10000, 11000], [11000, 10250, 13000], [9250, 9750, 11250], [11000, 10500, 14000]]
输出:
unique = [[11000, 10250, 13000], [11000, 10500, 14000]]
原因:索引0和2的数组彼此所有对应元素差值均在500以内,另外两个数组无法找到满足条件的其他数组。
原实现代码(低效):
import numpy as np def unique(intervals, k): unique = [] for i in range(len(intervals)): l = intervals[:i] + intervals[i+1:] t = intervals[i] n = 0 if ((np.array(l) >= np.array(t) - 500) & (np.array(l) <= np.array(t) + 500)).all(1).any(): n += 1 if n == 0: unique.append(intervals[i]) return unique
该代码在处理大数据集时速度极慢,核心问题在于:
- 每次循环都切片生成新列表并转换为numpy数组,产生大量临时对象,内存与时间开销大。
- 采用O(n²)的全量两两比较,数据量增大时性能呈指数级下降。
优化方案
方案1:全量向量化计算(适合中等规模数据集)
利用numpy的广播机制一次性完成所有两两数组的差值计算,避免循环中的重复操作。
import numpy as np def unique_optimized(intervals, k): arr = np.array(intervals) # 计算所有两两数组间的最大绝对差值(按维度取最大) max_diff = np.abs(arr[:, None] - arr).max(axis=2) # 填充对角线为501,排除自身与自身的比较 np.fill_diagonal(max_diff, 501) # 筛选出所有无匹配的数组索引(该行所有差值均>500) no_match_indices = np.where((max_diff > 500).all(axis=1))[0] return arr[no_match_indices].tolist()
核心逻辑:
- 将输入列表转为numpy数组,利用
arr[:, None] - arr广播生成(n,n,k)的差值数组,再取每个两两对的最大绝对差值。 - 若最大差值≤500,说明该对数组所有对应元素差值都在500以内。
- 排除自身比较后,筛选出没有任何匹配数组的结果。
方案2:分维度索引筛选(适合超大规模数据集)
通过对每个维度建立排序索引,快速缩小候选数组范围,减少不必要的全量比较。
import numpy as np from bisect import bisect_left, bisect_right def unique_optimized_with_index(intervals, k): arr = np.array(intervals) # 为每个维度建立排序后的数值-索引映射 sorted_dims = [] for j in range(k): dim_pairs = sorted(zip(arr[:, j], range(len(arr)))) sorted_dims.append(([x[0] for x in dim_pairs], [x[1] for x in dim_pairs])) unique_list = [] for idx, target in enumerate(arr): # 初始化候选集合,排除自身 candidates = set(range(len(arr))) candidates.remove(idx) # 按每个维度筛选候选,取交集 for j in range(k): val = target[j] sorted_vals, sorted_indices = sorted_dims[j] # 二分查找当前维度的数值范围边界 left_pos = bisect_left(sorted_vals, val - 500) right_pos = bisect_right(sorted_vals, val + 500) # 提取该范围内的数组索引,与候选集合取交集 dim_candidates = set(sorted_indices[left_pos:right_pos]) candidates.intersection_update(dim_candidates) if not candidates: break # 无候选则提前终止 # 验证候选数组是否满足全维度差值条件 has_match = False for cand_idx in candidates: if np.all(np.abs(arr[cand_idx] - target) <= 500): has_match = True break if not has_match: unique_list.append(target.tolist()) return unique_list
核心逻辑:
- 对每个维度的数值排序并保留原索引,利用二分查找快速定位满足差值范围的候选数组。
- 取所有维度候选的交集,仅对这些候选做全维度验证,大幅减少比较次数。
- 时间复杂度可从O(n²)降至O(nklogn + n*m)(m为单个数组的候选数量,远小于n)。
内容的提问来源于stack exchange,提问作者Ben Stewart
相关产品推荐
相关产品推荐

