基于numpy argpartition高效筛选k个目标元素并排除指定索引的优化
需求与优化诉求
我需要实现一种最高效的方法,从数组中筛选并排序出k个元素,同时要排除指定的部分索引。目前基于numpy的argpartition实现了一个函数,但该函数会处理多于实际需要的元素(为了后续过滤排除索引)。尝试过直接忽略指定索引进行排序但未成功,希望能借助argpartition的特性或其他方法优化性能。
当前实现代码
def find_cads_indices_closest_to_target_distance( cad_idx, embeddings, target_distance, closest_to_target_distance_count=1, excluded_indices=None, metric="cosine", ): # 默认将cad_idx加入排除索引 if excluded_indices is None: excluded_indices = [cad_idx] elif cad_idx not in excluded_indices: excluded_indices.append(cad_idx) reference_embedding = embeddings[cad_idx] # 根据指定度量计算距离 if metric == "cosine": distances = cosine_distances([reference_embedding], embeddings)[0] elif metric == "euclidean": distances = euclidean_distances([reference_embedding], embeddings)[0] else: raise ValueError("无效度量,请选择'cosine'或'euclidean'。") differences = np.abs(distances - target_distance) # 找到差异值最小的嵌入索引 indices_of_closest_to_target = np.argpartition( differences, closest_to_target_distance_count + len(excluded_indices) )[: closest_to_target_distance_count + len(excluded_indices)] # 从列表中排除指定索引 indices_of_closest_to_target = [ idx for idx in indices_of_closest_to_target if idx not in excluded_indices ] # 如果数量超过需求,截断列表 indices_of_closest_to_target = indices_of_closest_to_target[ :closest_to_target_distance_count ] # 可选:按距离排序 indices_of_closest_to_target = sorted( indices_of_closest_to_target, key=lambda idx: differences[idx] ) # 返回最接近的嵌入对应的索引及其距离 return ( indices_of_closest_to_target, distances[indices_of_closest_to_target], )
优化方案
核心思路:提前屏蔽排除索引,避免多余计算
直接将排除索引对应的差异值设为无穷大,这样argpartition会自动将这些索引排在末尾,无需额外多取元素再过滤,大幅减少后续处理步骤。
修改后的代码
import numpy as np from sklearn.metrics.pairwise import cosine_distances, euclidean_distances def find_cads_indices_closest_to_target_distance( cad_idx, embeddings, target_distance, closest_to_target_distance_count=1, excluded_indices=None, metric="cosine", ): # 默认将cad_idx加入排除索引,并转为集合加速操作 excluded = {cad_idx} if excluded_indices is not None: excluded.update(excluded_indices) reference_embedding = embeddings[cad_idx] # 根据指定度量计算距离 if metric == "cosine": distances = cosine_distances([reference_embedding], embeddings)[0] elif metric == "euclidean": distances = euclidean_distances([reference_embedding], embeddings)[0] else: raise ValueError("无效度量,请选择'cosine'或'euclidean'。") differences = np.abs(distances - target_distance) # 关键优化:将排除索引的差异值设为无穷大,确保它们不会被选中 if excluded: differences[list(excluded)] = np.inf # 直接取前k个最小差异值的索引 k = closest_to_target_distance_count top_k_indices = np.argpartition(differences, k)[:k] # 对前k个索引按差异值排序(若需要有序结果) top_k_indices = top_k_indices[np.argsort(differences[top_k_indices])] # 返回结果 return top_k_indices.tolist(), distances[top_k_indices].tolist()
优化点说明
- 屏蔽排除索引:通过将排除索引对应的
differences设为np.inf,argpartition会自动忽略这些索引,无需多取元素再过滤,节省了列表推导式的开销。 - 减少不必要操作:直接取前k个索引,避免了先取
k+len(excluded)个元素再截断的步骤,降低内存占用和计算量。 - 集合优化查找:将排除索引转为集合,初始化时的查找、更新操作时间复杂度从O(n)降为O(1)。
- numpy原生排序:用
np.argsort替代Python内置sorted,处理numpy数组的效率更高。
内容的提问来源于stack exchange,提问作者accoumar
相关产品推荐
相关产品推荐

