【发布时间】:2016-12-30 11:34:34
【问题描述】:
我刚刚注意到,我的脚本的执行时间几乎减半,只需将乘法更改为除法。
为了调查这个问题,我写了一个小例子:
import numpy as np
import timeit
# uint8 array
arr1 = np.random.randint(0, high=256, size=(100, 100), dtype=np.uint8)
# float32 array
arr2 = np.random.rand(100, 100).astype(np.float32)
arr2 *= 255.0
def arrmult(a):
"""
mult, read-write iterator
"""
b = a.copy()
for item in np.nditer(b, op_flags=["readwrite"]):
item[...] = (item + 5) * 0.5
def arrmult2(a):
"""
mult, index iterator
"""
b = a.copy()
for i, j in np.ndindex(b.shape):
b[i, j] = (b[i, j] + 5) * 0.5
def arrmult3(a):
"""
mult, vectorized
"""
b = a.copy()
b = (b + 5) * 0.5
def arrdiv(a):
"""
div, read-write iterator
"""
b = a.copy()
for item in np.nditer(b, op_flags=["readwrite"]):
item[...] = (item + 5) / 2
def arrdiv2(a):
"""
div, index iterator
"""
b = a.copy()
for i, j in np.ndindex(b.shape):
b[i, j] = (b[i, j] + 5) / 2
def arrdiv3(a):
"""
div, vectorized
"""
b = a.copy()
b = (b + 5) / 2
def print_time(name, t):
print("{: <10}: {: >6.4f}s".format(name, t))
timeit_iterations = 100
print("uint8 arrays")
print_time("arrmult", timeit.timeit("arrmult(arr1)", "from __main__ import arrmult, arr1", number=timeit_iterations))
print_time("arrmult2", timeit.timeit("arrmult2(arr1)", "from __main__ import arrmult2, arr1", number=timeit_iterations))
print_time("arrmult3", timeit.timeit("arrmult3(arr1)", "from __main__ import arrmult3, arr1", number=timeit_iterations))
print_time("arrdiv", timeit.timeit("arrdiv(arr1)", "from __main__ import arrdiv, arr1", number=timeit_iterations))
print_time("arrdiv2", timeit.timeit("arrdiv2(arr1)", "from __main__ import arrdiv2, arr1", number=timeit_iterations))
print_time("arrdiv3", timeit.timeit("arrdiv3(arr1)", "from __main__ import arrdiv3, arr1", number=timeit_iterations))
print("\nfloat32 arrays")
print_time("arrmult", timeit.timeit("arrmult(arr2)", "from __main__ import arrmult, arr2", number=timeit_iterations))
print_time("arrmult2", timeit.timeit("arrmult2(arr2)", "from __main__ import arrmult2, arr2", number=timeit_iterations))
print_time("arrmult3", timeit.timeit("arrmult3(arr2)", "from __main__ import arrmult3, arr2", number=timeit_iterations))
print_time("arrdiv", timeit.timeit("arrdiv(arr2)", "from __main__ import arrdiv, arr2", number=timeit_iterations))
print_time("arrdiv2", timeit.timeit("arrdiv2(arr2)", "from __main__ import arrdiv2, arr2", number=timeit_iterations))
print_time("arrdiv3", timeit.timeit("arrdiv3(arr2)", "from __main__ import arrdiv3, arr2", number=timeit_iterations))
这将打印以下时间:
uint8 arrays
arrmult : 2.2004s
arrmult2 : 3.0589s
arrmult3 : 0.0014s
arrdiv : 1.1540s
arrdiv2 : 2.0780s
arrdiv3 : 0.0027s
float32 arrays
arrmult : 1.2708s
arrmult2 : 2.4120s
arrmult3 : 0.0009s
arrdiv : 1.5771s
arrdiv2 : 2.3843s
arrdiv3 : 0.0009s
我一直认为乘法在计算上比除法便宜。然而,对于uint8,一个除法似乎几乎是两倍的效率。这是否与 * 0.5 必须计算浮点数中的乘法然后将结果转换回整数的事实有关?
至少对于浮点数乘法似乎比除法更快。这通常是真的吗?
为什么uint8 中的乘法比float32 中的更广泛?我认为 8 位无符号整数的计算速度应该比 32 位浮点数快得多?!
有人可以“揭秘”这个吗?
编辑:为了获得更多数据,我加入了矢量化函数(如建议的那样)并添加了索引迭代器。矢量化函数要快得多,因此不能真正具有可比性。然而,如果 timeit_iterations 为向量化函数设置得更高,则结果表明,uint8 和 float32 两者的乘法都更快。我想这更令人困惑?!
也许乘法实际上总是比除法快,但 for 循环中的主要性能漏洞不是算术运算,而是循环本身。虽然这并不能解释为什么循环在不同的操作中表现不同。
EDIT2:就像@jotasi 已经说过的那样,我们正在寻找division 与multiplication 和int(或uint8)与float 的完整解释(或float32)。此外,解释向量化方法和迭代器的不同趋势会很有趣,因为在向量化情况下,除法似乎较慢,而在迭代器情况下则更快。
【问题讨论】:
-
如果我不忽略一些愚蠢的事情,它会变得更加陌生。用
b = (b+5) * 0.5和b = (b+5) / 2替换for 循环会导致除法变慢。 -
顺便说一句,您的代码中有一个小错字。您应该将
arrdiv3中的b = (b + 5) / 0.5更改为b = (b + 5) / 2。 -
你是对的!谢谢,我已经修好了。
-
我也调整了时间。奇怪的是
div3在float32中的时间没有改变,而uint8上升到2.7E-3。我猜矢量化版本的时间太短了,无法提供精确的测量值?! -
可能是这样。但我也用更大的数组检查了它们并得到了相同的趋势(除法更慢)。也许您可以强调,对矢量化与迭代器以及 int 与 float 的完整解释会很好。
标签: python performance python-2.7 numpy