【问题标题】:how does numpy.where work?numpy.where 是如何工作的?
【发布时间】:2017-02-01 06:25:41
【问题描述】:

我可以理解以下 numpy 行为。

>>> a
array([[ 0. ,  0. ,  0. ],
       [ 0. ,  0.7,  0. ],
       [ 0. ,  0.3,  0.5],
       [ 0.6,  0. ,  0.8],
       [ 0.7,  0. ,  0. ]])
>>> argmax_overlaps = a.argmax(axis=1)
>>> argmax_overlaps
array([0, 1, 2, 2, 0])
>>> max_overlaps = a[np.arange(5),argmax_overlaps]
>>> max_overlaps
array([ 0. ,  0.7,  0.5,  0.8,  0.7])
>>> gt_argmax_overlaps = a.argmax(axis=0)
>>> gt_argmax_overlaps
array([4, 1, 3])
>>> gt_max_overlaps = a[gt_argmax_overlaps,np.arange(a.shape[1])]
>>> gt_max_overlaps
array([ 0.7,  0.7,  0.8])
>>> gt_argmax_overlaps = np.where(a == gt_max_overlaps)
>>> gt_argmax_overlaps
(array([1, 3, 4]), array([1, 2, 0]))

我知道 0.7, 0.7 和 0.8 是 a[1,1],a[3,2] 和 a[4,0] 所以我得到了元组 (array[1,3,4] and array[1,2,0]) 每个数组由第 0 个和第一个索引组成这三个要素。然后我尝试了其他示例,以查看我的理解是否正确。

>>> np.where(a == [0.3])
(array([2]), array([1]))

0.3 在 a[2,1] 中,所以结果看起来和我预期的一样。然后我尝试了

>>> np.where(a == [0.3, 0.5])
(array([], dtype=int64),)

??我希望看到 (array([2,2]),array([2,3]))。为什么我会看到上面的输出?

>>> np.where(a == [0.7, 0.7, 0.8])
(array([1, 3, 4]), array([1, 2, 0]))
>>> np.where(a == [0.8,0.7,0.7])
(array([1]), array([1]))

我也无法理解第二个结果。有人可以向我解释一下吗?谢谢。

【问题讨论】:

  • 使用np.where((a==0.3)|(a==0.5))np.where((a==0.7)|(a==0.8)) 获得正确的结果。但是我不知道为什么np.where(a == [0.7, 0.7, 0.8]) 有效,而np.where(a==[0.7,0.8]) 抛出DeprecationWarning。看起来像一个错误。
  • where 给出意外索引时,查看条件数组。 where 只是告诉你该数组在哪里True

标签: python numpy where


【解决方案1】:

首先要意识到np.where(a == [whatever]) 只是向您显示a == [whatever] 为True 的索引。因此,您可以通过查看a == [whatever] 的值来获得提示。在您的情况下“有效”:

>>> a == [0.7, 0.7, 0.8]
array([[False, False, False],
       [False,  True, False],
       [False, False, False],
       [False, False,  True],
       [ True, False, False]], dtype=bool)

你没有得到你认为的那样。您认为这是分别要求每个元素的索引,但它获取的是值匹配的位置在行中的相同位置。基本上这个比较是在说“对于每一行,告诉我第一个元素是否为 0.7,第二个元素是否为 0.7,第三个元素是否为 0.8”。然后它返回那些匹配位置的索引。换句话说,比较是在整行之间进行的,而不仅仅是单个值。对于你的最后一个例子:

>>> a == [0.8,0.7,0.7]
array([[False, False, False],
       [False,  True, False],
       [False, False, False],
       [False, False, False],
       [False, False, False]], dtype=bool)

您现在得到不同的结果。它不是要求“a 的值为 0.8 的索引”,它只要求在行的开头有 0.8 的索引 - 同样在任何一个中都有一个 0.7后面两个位置。

只有当您比较的值与a 的单行形状匹配时,才能进行这种逐行比较。因此,当您尝试使用二元素列表时,它会返回一个空集,因为它试图将列表作为标量值与数组中的各个值进行比较。

结果是您不能在值列表上使用== 并期望它只告诉您任何值出现的位置。相等将按值 和位置 匹配(如果您比较的值与数组中的一行形状相同),或者它将尝试将整个列表作为标量进行比较(如果形状不匹配)。如果您想独立搜索值,则需要执行 Khris 在评论中建议的操作:

np.where((a==0.3)|(a==0.5))

也就是说,您需要对单独的值进行两次(或更多)单独的比较,而不是对值列表进行一次比较。

【讨论】:

  • 哇,原来如此。 python既聪明又奇怪:)
猜你喜欢
  • 2011-08-04
  • 1970-01-01
  • 1970-01-01
  • 2012-03-19
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2018-11-11
相关资源
最近更新 更多