You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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()

优化点说明

  1. 屏蔽排除索引:通过将排除索引对应的differences设为np.inf,argpartition会自动忽略这些索引,无需多取元素再过滤,节省了列表推导式的开销。
  2. 减少不必要操作:直接取前k个索引,避免了先取k+len(excluded)个元素再截断的步骤,降低内存占用和计算量。
  3. 集合优化查找:将排除索引转为集合,初始化时的查找、更新操作时间复杂度从O(n)降为O(1)。
  4. numpy原生排序:用np.argsort替代Python内置sorted,处理numpy数组的效率更高。

内容的提问来源于stack exchange,提问作者accoumar

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 12:53:12