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

求替代numpy.where()的方法:按指定值顺序返回索引而非升序

问题:按指定值顺序获取对应索引,替代numpy.where()的高效方法

我用numpy.where()从列表中获取在另一列表中出现的值的索引,但它始终返回升序排列的索引列表。示例代码如下:

import numpy as np
lst = [1, 2, 8, 7, 3, 4, 6, 5]
values = [4, 8]
indices = np.where(np.isin(lst, values))

输出结果为[2,5],但我期望得到[5,2]——按values中值的顺序返回对应索引,先返回4对应的索引5,再返回8对应的索引2。

这是神经网络自定义2D最大池化函数的一部分,速度至关重要,不想用循环实现,请问有没有替代numpy.where()的方法?


解决方案

可以用值到索引的矢量化映射来实现,全程无Python循环,完全依赖numpy的底层矢量化操作,速度符合神经网络场景的要求:

import numpy as np

lst = [1, 2, 8, 7, 3, 4, 6, 5]
values = [4, 8]

# 创建值到对应索引的数组映射(假设lst中的值都是非负整数,且最大值不大)
max_val = np.max(lst)
value_to_idx = np.zeros(max_val + 1, dtype=np.int64)
value_to_idx[lst] = np.arange(len(lst))

# 直接按values的顺序提取索引
indices = value_to_idx[values]
print(indices)  # 输出: [5 2]

补充说明:

  • 这个方法的核心是利用numpy的数组索引特性,一次性完成所有值到索引的映射,时间复杂度为O(n),比循环高效得多。
  • 如果lst中存在重复值,上述代码会保留该值最后一次出现的索引;如果需要获取所有出现的索引,可以结合np.where和列表推导(重复值较少时性能仍可接受):
    # 处理重复值的场景
    indices = [np.where(np.array(lst) == val)[0] for val in values]
    # 输出会是数组的列表,比如若lst中有两个4,会返回[[5, x], [2]]
    
  • 若lst中的值范围很大(比如超过1e6),用数组映射会浪费内存,此时可以改用np.unique创建字典映射:
    unique_vals, unique_indices = np.unique(lst, return_index=True)
    value_to_idx = dict(zip(unique_vals, unique_indices))
    indices = np.array([value_to_idx[val] for val in values])
    
    这种方式内存占用更小,同样是矢量化为主的操作,性能损失可以忽略。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 16:51:13