求替代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
相关产品推荐
相关产品推荐

