【问题标题】:testing whether a Numpy array contains a given row测试 Numpy 数组是否包含给定的行
【发布时间】:2021-02-27 03:47:57
【问题描述】:

是否有一种 Pythonic 和有效的方法来检查 Numpy 数组是否包含给定行的至少一个实例? “高效”是指它在找到第一个匹配行时终止,而不是遍历整个数组,即使已经找到结果。

对于 Python 数组,这可以通过 if row in array: 非常干净地完成,但这并不像我对 Numpy 数组所期望的那样工作,如下所示。

使用 Python 数组:

>>> a = [[1,2],[10,20],[100,200]]
>>> [1,2] in a
True
>>> [1,20] in a
False

但是 Numpy 数组给出了不同且看起来很奇怪的结果。 (ndarray__contains__ 方法似乎没有记录。)

>>> a = np.array([[1,2],[10,20],[100,200]])
>>> np.array([1,2]) in a
True
>>> np.array([1,20]) in a
True
>>> np.array([1,42]) in a
True
>>> np.array([42,1]) in a
False

【问题讨论】:

  • 您希望实现不可能。 Numpy 目前不提供任何会在找到第一个时停止的东西。但是,如果您这样做更多,那么几次基于排序的方法无论如何都会更有效。至于__contains__ 的行为,我几乎可以说这是一个错误(即它适用于标量,但数组有点奇怪,尽管在内部它只是 tom10 说的)
  • @seberg 你确定不存在解决方案吗?如果是这样,那么这就是我的问题的答案,所以请发布它,如果我确信我会接受它。此外,如果您可以通过“基于排序的方法”来解释您的意思,那将会很有帮助。我的数组实际上是排序的,因此最常搜索的行往往靠近顶部,如果这就是你的意思的话 - 但这没有用,除非查询方法在找到匹配项后停止。
  • @seberg 如果__collect__ 正在执行tom10 所说的操作,那么我的问题中引用的最后一行输入将返回True,不是吗?

标签: python numpy


【解决方案1】:

你可以使用 .tolist()

>>> a = np.array([[1,2],[10,20],[100,200]])
>>> [1,2] in a.tolist()
True
>>> [1,20] in a.tolist()
False
>>> [1,20] in a.tolist()
False
>>> [1,42] in a.tolist()
False
>>> [42,1] in a.tolist()
False

或者使用视图:

>>> any((a[:]==[1,2]).all(1))
True
>>> any((a[:]==[1,20]).all(1))
False

或者通过 numpy 列表生成(可能非常慢):

any(([1,2] == x).all() for x in a)     # stops on first occurrence 

或者使用numpy逻辑函数:

any(np.equal(a,[1,2]).all(1))

如果你计时:

import numpy as np
import time

n=300000
a=np.arange(n*3).reshape(n,3)
b=a.tolist()

t1,t2,t3=a[n//100][0],a[n//2][0],a[-10][0]

tests=[ ('early hit',[t1, t1+1, t1+2]),
        ('middle hit',[t2,t2+1,t2+2]),
        ('late hit', [t3,t3+1,t3+2]),
        ('miss',[0,2,0])]

fmt='\t{:20}{:.5f} seconds and is {}'     

for test, tgt in tests:
    print('\n{}: {} in {:,} elements:'.format(test,tgt,n))

    name='view'
    t1=time.time()
    result=(a[...]==tgt).all(1).any()
    t2=time.time()
    print(fmt.format(name,t2-t1,result))

    name='python list'
    t1=time.time()
    result = True if tgt in b else False
    t2=time.time()
    print(fmt.format(name,t2-t1,result))

    name='gen over numpy'
    t1=time.time()
    result=any((tgt == x).all() for x in a)
    t2=time.time()
    print(fmt.format(name,t2-t1,result))

    name='logic equal'
    t1=time.time()
    np.equal(a,tgt).all(1).any()
    t2=time.time()
    print(fmt.format(name,t2-t1,result))

可以看到hit or miss,numpy 例程搜索数组的速度是一样的。 Python in 运算符可能在早期命中时要快得多,如果你必须一直遍历数组,生成器只是个坏消息。

以下是 300,000 x 3 元素数组的结果:

early hit: [9000, 9001, 9002] in 300,000 elements:
    view                0.01002 seconds and is True
    python list         0.00305 seconds and is True
    gen over numpy      0.06470 seconds and is True
    logic equal         0.00909 seconds and is True

middle hit: [450000, 450001, 450002] in 300,000 elements:
    view                0.00915 seconds and is True
    python list         0.15458 seconds and is True
    gen over numpy      3.24386 seconds and is True
    logic equal         0.00937 seconds and is True

late hit: [899970, 899971, 899972] in 300,000 elements:
    view                0.00936 seconds and is True
    python list         0.30604 seconds and is True
    gen over numpy      6.47660 seconds and is True
    logic equal         0.00965 seconds and is True

miss: [0, 2, 0] in 300,000 elements:
    view                0.00936 seconds and is False
    python list         0.01287 seconds and is False
    gen over numpy      6.49190 seconds and is False
    logic equal         0.00965 seconds and is False

对于 3,000,000 x 3 数组:

early hit: [90000, 90001, 90002] in 3,000,000 elements:
    view                0.10128 seconds and is True
    python list         0.02982 seconds and is True
    gen over numpy      0.66057 seconds and is True
    logic equal         0.09128 seconds and is True

middle hit: [4500000, 4500001, 4500002] in 3,000,000 elements:
    view                0.09331 seconds and is True
    python list         1.48180 seconds and is True
    gen over numpy      32.69874 seconds and is True
    logic equal         0.09438 seconds and is True

late hit: [8999970, 8999971, 8999972] in 3,000,000 elements:
    view                0.09868 seconds and is True
    python list         3.01236 seconds and is True
    gen over numpy      65.15087 seconds and is True
    logic equal         0.09591 seconds and is True

miss: [0, 2, 0] in 3,000,000 elements:
    view                0.09588 seconds and is False
    python list         0.12904 seconds and is False
    gen over numpy      64.46789 seconds and is False
    logic equal         0.09671 seconds and is False

这似乎表明np.equal 是最快的纯 numpy 方法...

【讨论】:

  • 谢谢,但我正在寻找一种在找到第一个匹配行后终止的实现,而不是像tolist 那样遍历整个数组。问题的第一个版本对此并不清楚;我已经编辑过了。
  • view 方法是否懒惰地评估?我怀疑在视图上调用 .all 会创建一个全新的数组,但我不知道如何找到。
  • 几点:1)在“视图”中,我认为你应该使用a[...],而不是a[:]; 2)在“逻辑”中,我认为你应该使用 np.any 和 np.all 而不是 python 的; 3) 对 False 结果进行比较也是很好的,因为对于其中一些情况(尤其是“gen”),这将有很大不同。 +1,不过,对于实际的效率衡量标准。
  • @Pyson 非常感谢您的更新。我可以看到,在我的用例中, np.equal 可能会比使用 Python 列表更快,即使它没有因提前终止而获得奖励。了解这一点非常有用。
  • 定时结果最好使用timeit 模块:使用time.time() 方法,如果系统在后台运行另一个任务(报告的时间太大),结果可能会非常不准确。
【解决方案2】:

Numpys __contains__ is, at the time of writing this, (a == b).any() 只有当 b 是一个标量时才可能是正确的(它有点毛茸茸,但我相信 - 仅在 1.7 或更高版本中像这样工作 - 这将是正确的通用方法(a == b).all(np.arange(a.ndim - b.ndim, a.ndim)).any(),这对ab 维度的所有组合都有意义)...

编辑:为了清楚起见,这不一定涉及广播时的预期结果。也有人可能会争辩说它应该像np.in1d 那样单独处理a 中的项目。我不确定它应该有一种明确的工作方式。

现在您希望 numpy 在找到第一次出现时停止。这个 AFAIK 目前不存在。这很困难,因为 numpy 主要基于 ufunc,它们在整个数组上做同样的事情。 Numpy 确实优化了这类归约,但只有在被归约的数组已经是布尔数组(即np.ones(10, dtype=bool).any())时才有效。

否则它需要一个不存在的__contains__ 的特殊函数。这可能看起来很奇怪,但您必须记住,numpy 支持许多数据类型,并且具有更大的机制来选择正确的数据类型并选择正确的函数来处理它。因此,换句话说,ufunc 机器无法做到这一点,并且由于数据类型的原因,实现__contains__ 或类似的东西实际上并不是那么简单。

你当然可以用python写,或者你可能知道你的数据类型,用Cython/C自己写很简单。


那是说。通常,对这些事情使用基于排序的方法要好得多。这有点乏味,而且对于 lexsort 没有 searchsorted 这样的东西,但它有效(如果你愿意,你也可以滥用 scipy.spatial.cKDTree)。这假设您只想沿最后一个轴进行比较:

# Unfortunatly you need to use structured arrays:
sorted = np.ascontiguousarray(a).view([('', a.dtype)] * a.shape[-1]).ravel()

# Actually at this point, you can also use np.in1d, if you already have many b
# then that is even better.

sorted.sort()

b_comp = np.ascontiguousarray(b).view(sorted.dtype)
ind = sorted.searchsorted(b_comp)

result = sorted[ind] == b_comp

这也适用于数组b,如果您保留已排序的数组,那么在b 中的单个值(行)一次执行此操作也会更好,此时a 保持不变一样的(否则我只会np.in1d 在将其视为recarray 之后)。 重要提示:为了安全起见,您必须执行np.ascontiguousarray。它通常什么都不做,但如果有,那将是一个很大的潜在错误。

【讨论】:

  • 谢谢,这很有帮助。我会等几天,以防有人知道一些特别聪明的解决方案,如果没有,我会接受这个答案。 (显然,我只是对 (a==b).any() 将返回的内容有点密集。)
  • ind == len(sorted) 时出现 IndexError。如果b“超出”sorted 数组,就会发生这种情况;例如b = [101,0].
【解决方案3】:

我认为

equal([1,2], a).all(axis=1)   # also,  ([1,2]==a).all(axis=1)
# array([ True, False, False], dtype=bool)

将列出匹配的行。正如 Jamie 指出的那样,要知道是否存在至少一个这样的行,请使用 any

equal([1,2], a).all(axis=1).any()
# True

除此之外:
我怀疑in(和__contains__)和上面一样,但使用any而不是all

【讨论】:

  • +1 不错!不过,您必须将其全部包含在 np.any(...) 中才能获得一个成员资格布尔值。
  • 谢谢。但这将遍历整个数组并在内存中分配一个包含所有结果的新数组,然后才检查它是否为空。一个高效的实现会在找到第一个匹配行后立即停止并返回 True。
  • 我已经编辑了这个问题,以澄清我所说的“高效”。
【解决方案4】:

我将建议的解决方案与 perfplot 进行了比较,发现如果您要在一个长的未排序列表中寻找一个 2 元组,

np.any(np.all(a == b, axis=1))

是最快的解决方案。如果在前几行中找到匹配项,则显式短路循环总是会更快。

重现情节的代码:

import numpy as np
import perfplot


target = [6, 23]


def setup(n):
    return np.random.randint(0, 100, (n, 2))


def any_all(data):
    return np.any(np.all(target == data, axis=1))


def tolist(data):
    return target in data.tolist()

def loop(data):
    for row in data:
        if np.all(row == target):
            return True
    return False


def searchsorted(a):
    s = np.ascontiguousarray(a).view([('', a.dtype)] * a.shape[-1]).ravel()
    s.sort()
    t = np.ascontiguousarray(target).view(s.dtype)
    ind = s.searchsorted(t)
    return (s[ind] == t)[0]


perfplot.save(
    "out02.png",
    setup=setup,
    kernels=[any_all, tolist, loop, searchsorted],
    n_range=[2 ** k for k in range(2, 20)],
    xlabel="len(array)",
)

【讨论】:

    【解决方案5】:

    如果你真的想在第一次出现时停止,你可以写一个循环,比如:

    import numpy as np
    
    needle = np.array([10, 20])
    haystack = np.array([[1,2],[10,20],[100,200]])
    found = False
    for row in haystack:
        if np.all(row == needle):
            found = True
            break
    print("Found: ", found)
    

    但是,我强烈怀疑,它会比使用 numpy 例程为整个数组执行此操作的其他建议慢得多。

    【讨论】:

    • 是的,这就是我试图避免的。如果事实证明 Numpy 不提供这样做的内置方式,那将是令人失望的。
    • 老实说,如果你知道它通常是在数组的乞求时,这不是一个糟糕的解决方案(如果needle 是一个足够大的数组,那么无论如何它都是一个很好的解决方案)。
    • any(np.all(row == needle) for row in haystack) 似乎更 Pythonic,而不是使用布尔标志。这个still short circuits.
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2013-08-26
    • 2010-11-13
    相关资源
    最近更新 更多