Numpy:如何用where()获取元组数组中指定元组的索引
关于numpy.where获取元组数组索引的疑问解答
你对
numpy.all(axis=1)的理解完全正确:
执行a == (1, 2)得到的是2维布尔数组,numpy.all(axis=1)会对每行元素做逻辑与运算,最终生成形状为(3,)的一维布尔数组,对应结果是array([False, True, False]),这一步的行为符合预期。为什么
numpy.where的返回需要用b[0][0]才能拿到目标索引?
因为numpy.where的返回值是一个元组,元组中的每个元素对应输入数组某一维度的匹配索引集合。针对一维的布尔数组输入,返回的元组里仅包含一个元素——存储匹配索引的一维数组,所以你得到的b实际结构是(array([1]),),而非单独的数组。
因此必须先通过b[0]取出这个一维索引数组,再用[0]提取其中的单个索引值。简化索引获取的写法:
如果你能确定目标元组在数组中仅有一个匹配项,可以用更简洁的方式获取索引:import numpy as np a = np.array([(0, 1), (1, 2), (2, 3)]) # 方式1:使用argwhere idx = np.argwhere(np.all(a == (1, 2), axis=1))[0][0] # 方式2:使用flatnonzero idx = np.flatnonzero(np.all(a == (1, 2), axis=1))[0]
内容的提问来源于stack exchange,提问作者Thomas Slade
相关产品推荐
相关产品推荐

