方法#1
获取与您正在寻找的数组的相等性,给我们一个2D 数组。然后,查找所有与 .all(axis=1) 匹配的行,这将是一个 1D 布尔数组。最后,要获取匹配项中的第一个实例,请使用.argmax() 并沿从开始到该索引的行对数组进行切片。
因此,完整的实现将是 -
s[:(s == [0,0,0,1]).all(1).argmax()]
示例逐步运行 -
In [39]: s # Input array
Out[39]:
array([[1, 0, 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1],
[0, 0, 0, 1],
[0, 1, 0, 0]])
In [33]: s == [0,0,0,1] # compare against search array
Out[33]:
array([[False, True, True, False],
[ True, True, False, False],
[ True, True, True, True],
[ True, True, True, True],
[ True, False, True, False]], dtype=bool)
In [34]: (s == [0,0,0,1]).all(1)
Out[34]: array([False, False, True, True, False], dtype=bool)
In [37]: (s == [0,0,0,1]).all(1).argmax()
Out[37]: 2
In [38]: s[:(s == [0,0,0,1]).all(1).argmax()]
Out[38]:
array([[1, 0, 0, 0],
[0, 0, 1, 0]])
方法 #2
由于我们处理的是单热编码数组,我们可以在2D 输入数组的每一行使用argmax,从而将其减少为1D 数组。同样,将搜索数组减少为标量,其余步骤保持不变。这将提高内存效率,因为我们将避免创建 2D 布尔数组。让我们直接进入示例运行 -
In [89]: s.argmax(1)
Out[89]: array([0, 2, 3, 3, 1])
In [90]: np.argmax([0,0,0,1])
Out[90]: 3
In [91]: s.argmax(1) == np.argmax([0,0,0,1])
Out[91]: array([False, False, True, True, False], dtype=bool)
In [92]: (s.argmax(1) == np.argmax([0,0,0,1])).argmax()
Out[92]: 2
# Final code
In [93]: s[:(s.argmax(1) == np.argmax([0,0,0,1])).argmax()]
Out[93]:
array([[1, 0, 0, 0],
[0, 0, 1, 0]])