【问题标题】:Why "numpy.any" has no short-circuit mechanism?为什么“numpy.any”没有短路机制?
【发布时间】:2021-09-26 03:55:06
【问题描述】:

我不明白为什么还没有完成如此基本的优化:

In [1]: one_million_ones = np.ones(10**6)
In [2]: %timeit one_million_ones.any()
100 loops, best of 3: 693µs per loop

In [3]: ten_millions_ones = np.ones(10**7)
In [4]: %timeit ten_millions_ones.any()
10 loops, best of 3: 7.03 ms per loop

扫描整个数组,即使结论是第一项的证据。

【问题讨论】:

标签: python performance numpy


【解决方案1】:

短路是有代价的。您需要在代码中引入分支。

分支(例如if 语句)的问题在于它们可能比使用替代操作(没有分支)要慢,然后您还有可能包含大量开销的分支预测。

还取决于编译器和处理器,无分支代码可以使用处理器矢量化。我不是这方面的专家,但可能是某种SIMD 或 SSE?

我将在这里使用 numba,因为代码易于阅读且速度足够快,因此性能会根据这些小差异而发生变化:

import numba as nb
import numpy as np

@nb.njit
def any_sc(arr):
    for item in arr:
        if item:
            return True
    return False

@nb.njit
def any_not_sc(arr):
    res = False
    for item in arr:
        res |= item
    return res

arr = np.zeros(100000, dtype=bool)
assert any_sc(arr) == any_not_sc(arr)
%timeit any_sc(arr)
# 126 µs ± 7.12 µs per loop (mean ± std. dev. of 7 runs, 10000 loops each)
%timeit any_not_sc(arr)
# 15.5 µs ± 962 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)
%timeit arr.any()
# 31.1 µs ± 184 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)

在没有分支的最坏情况下,它几乎快 10 倍。但在最好的情况下,短路功能要快得多:

arr = np.zeros(100000, dtype=bool)
arr[0] = True
%timeit any_sc(arr)
# 1.97 µs ± 12.9 ns per loop (mean ± std. dev. of 7 runs, 1000000 loops each)
%timeit any_not_sc(arr)
# 15.1 µs ± 368 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)
%timeit arr.any()
# 31.2 µs ± 2.23 µs per loop (mean ± std. dev. of 7 runs, 10000 loops each)

所以这是一个应该优化哪种情况的问题:最好的情况?最坏的情况?平均情况(any 的平均情况是多少)?

这可能是 NumPy 开发人员想要优化最坏情况而不是最佳情况。还是他们根本不在乎?或者,也许他们只是想要“可预测”的性能。


请注意您的代码:您测量创建数组所需的时间以及执行any 所需的时间。如果any 发生短路,您的代码就不会注意到它!

%timeit np.ones(10**6)
# 9.12 ms ± 635 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit np.ones(10**7)
# 86.2 ms ± 5.15 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)

对于支持您的问题的决定性时间,您应该改用这个:

arr1 = np.ones(10**6)
arr2 = np.ones(10**7)
%timeit arr1.any()
# 4.04 ms ± 121 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit arr2.any()
# 39.8 ms ± 1.34 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)

【讨论】:

  • 感谢您详细的回答。
  • 我怀疑 Numba 生成的代码效率低下会影响您的时间安排。短路不应该在没有回报时产生那种灾难性的影响。额外的分支很容易预测。
  • @user2357112 是的,看起来太多了,但是分支总是有开销,因为即使预测总是正确的,它仍然需要在某个时候“检查”。 numba 也有可能意识到无分支的可以使用处理器矢量化,并且在第一种情况下甚至没有尝试它们。我没有时间研究我的例子中的特殊性。我怀疑通过一些专门的努力并直接在 C 中进行编码 - 最坏情况下的时间差会更低(可能只是 2 倍或更小),但在最坏情况下分支代码会更慢。
  • 是的,但问题是,循环实际上不必等待 进行检查。检查可以与循环继续其工作并行发生。我认为这些天正确预测的分支可能实际上是零延迟。
  • @user2357112 我真的不确定。我刚刚用 cython 进行了尝试:在最坏的情况下,这两个函数的速度大致相同,但让我感到奇怪的是,两者的速度几乎与短路 numba 函数一样快。我怀疑 numba 对于短路情况可能不是“效率低下”,但在非短路功能方面可能非常有效。但是,我现在真的没有时间真正检查 numba 的汇编或 cython 的代码。也许在周末之后。
【解决方案2】:

这是一个不固定的性能回归。 NumPy issue 3446. 实际上 short-circuiting logic,但是对ufunc.reduce 机制的更改引入了围绕短路逻辑的不必要的基于块的外循环,并且该外循环不知道如何短路。大家可以看看分块机制here的一些解释。

即使没有回归,短路效应也不会出现在您的测试中。首先,您正在为数组创建计时,其次,我认为他们从未为任何输入 dtype 放入短路逻辑,但布尔值。从讨论中,听起来numpy.any 背后的 ufunc 缩减机制的细节会让这变得困难。

讨论确实提出了一个令人惊讶的点,即argminargmax 方法似乎对布尔输入短路。 A quick test 表明,从 NumPy 1.12 开始(不是最新版本,而是 Ideone 上当前的版本),x[x.argmax()] 短路,它在 1 维布尔输入方面胜过 x.any()x.max()无论输入是小还是大,无论短路是否有效。奇怪!

【讨论】:

    猜你喜欢
    • 2016-10-24
    • 2011-07-09
    • 2019-05-25
    • 2013-06-12
    • 2010-12-17
    • 2014-06-14
    • 1970-01-01
    相关资源
    最近更新 更多