【问题标题】:NumPy performance: uint8 vs. float and multiplication vs. division?NumPy 性能:uint8 与浮点数和乘法与除法?
【发布时间】: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 为向量化函数设置得更高,则结果表明,uint8float32 两者的乘法都更快。我想这更令人困惑?!

也许乘法实际上总是比除法快,但 for 循环中的主要性能漏洞不是算术运算,而是循环本身。虽然这并不能解释为什么循环在不同的操作中表现不同。

EDIT2:就像@jotasi 已经说过的那样,我们正在寻找divisionmultiplicationint(或uint8)与float 的完整解释(或float32)。此外,解释向量化方法和迭代器的不同趋势会很有趣,因为在向量化情况下,除法似乎较慢,而在迭代器情况下则更快。

【问题讨论】:

  • 如果我不忽略一些愚蠢的事情,它会变得更加陌生。用b = (b+5) * 0.5b = (b+5) / 2 替换for 循环会导致除法变慢。
  • 顺便说一句,您的代码中有一个小错字。您应该将arrdiv3 中的b = (b + 5) / 0.5 更改为b = (b + 5) / 2
  • 你是对的!谢谢,我已经修好了。
  • 我也调整了时间。奇怪的是div3float32 中的时间没有改变,而uint8 上升到2.7E-3。我猜矢量化版本的时间太短了,无法提供精确的测量值?!
  • 可能是这样。但我也用更大的数组检查了它们并得到了相同的趋势(除法更慢)。也许您可以强调,对矢量化与迭代器以及 int 与 float 的完整解释会很好。

标签: python performance python-2.7 numpy


【解决方案1】:

这个答案只关注向量化操作,因为ead 已经回答了其他操作缓慢的原因。

很多“优化”都是基于旧硬件。意味着优化在旧硬件上成立的假设在新硬件上并不成立。

管道和划分

除法 很慢。除法运算由几个单元组成,每个单元必须一个接一个地执行一个计算。这就是使分裂缓慢的原因。

但是,在浮点处理单元 (FPU) [在大多数现代 CPU 上很常见] 中,有专门的单元排列在除法指令的“管道”中。一旦一个单元完成,该单元就不再需要用于其余的操作。如果您有多个划分操作,您可以在下一次划分操作时让这些单元无事可做。所以虽然每次操作都很慢,但FPU实际上可以实现除法运算的高吞吐量。流水线化与矢量化不同,但结果基本相同——当您有很多相同的操作要做时,吞吐量会更高。

将流水线视为流量。比较以每小时 30 英里的速度行驶的三条车道与以每小时 90 英里的速度行驶的一条车道。较慢的车流肯定是个别较慢,但三车道的车流量还是一样的。

【讨论】:

    【解决方案2】:

    问题在于您的假设,即您测量除法或乘法所需的时间,这是不正确的。您正在测量除法或乘法所需的开销。

    人们确实必须查看确切的代码来解释每种效果,这些效果可能因版本而异。这个答案只能给出一个想法,必须考虑什么。

    问题是一个简单的int 在 python 中一点都不简单:它是一个必须在垃圾收集器中注册的真实对象,它的大小随着它的价值而增长——你必须付出的一切:例如,对于 8 位整数,需要 24 字节的内存! python-floats 也类似。

    另一方面,numpy 数组由简单的 c 样式整数/浮点数组成,没有开销,您可以节省大量内存,但在访问 numpy-array 的元素期间会为此付出代价。 a[i] 意味着:必须构造一个 python 整数,在垃圾收集器中注册,并且只能使用它 - 有很多开销。

    考虑这段代码:

    li1=[x%256 for x in xrange(10**4)]
    arr1=np.array(li1, np.uint8)
    
    def arrmult(a):    
        for i in xrange(len(a)):
            a[i]*=5;
    

    arrmult(li1)arrmult(arr1) 快 25,因为列表中的整数已经是 python-ints 并且不必创建!创建对象需要大部分计算时间 - 几乎可以忽略其他所有内容。


    让我们看看你的代码,首先是乘法:

    def arrmult2(a):
        ...
        b[i, j] = (b[i, j] + 5) * 0.5
    

    在 uint8 的情况下,必须发生以下情况(为简单起见,我忽略了 +5):

    1. 必须创建一个 python-int
    2. 必须将其转换为浮点数(python-float 创建),才能进行浮点数乘法
    3. 并转换回 python-int 或/和 uint8

    对于 float32,要做的工作更少(乘法成本不高): 1. 创建了一个 python-float 2. 投回 float32。

    所以浮动版本应该更快,而且确实如此。


    现在我们来看看划分:

    def arrdiv2(a):
        ...
        b[i, j] = (b[i, j] + 5)  / 2 
    

    这里的陷阱:所有操作都是整数操作。因此,与乘法相比,无需转换为 python-float,因此与乘法相比,我们的开销更少。在您的情况下,unint8 的除法比乘法“更快”。

    但是,对于 float32,除法和乘法同样快/慢,因为在这种情况下几乎没有任何变化 - 我们仍然需要创建一个 python-float。


    现在是矢量化版本:它们与 c 风格的“原始”float32s/uint8s 一起工作,无需转换(及其成本!)到引擎盖下的相应 python 对象。要获得有意义的结果,您应该增加迭代次数(现在运行时间太短,无法肯定地说)。

    1. float32 的除法和乘法可能具有相同的运行时间,因为我希望 numpy 通过乘以 0.5 来替换除以 2(但要确保必须查看代码)。

    2. uint8 的乘法应该更慢,因为每个 uint8 整数必须在乘以 0.5 之前转换为浮点数,然后再转换回 uint8。

    3. 对于 uint8 的情况,numpy 不能通过乘以 0.5 来代替除以 2,因为它是整数除法。对于许多架构,整数除法比浮点乘法慢 - 这是最慢的矢量化操作。


    PS:我不会过多地谈论成本乘法与除法 - 有太多其他因素会对性能产生更大的影响。例如创建不必要的临时对象,或者如果 numpy-array 很大并且不适合缓存,那么内存访问将成为瓶颈 - 你会发现乘法和除法之间根本没有区别。

    【讨论】:

      【解决方案3】:

      这是因为您将 int 乘以 float 并将结果存储为 int。 尝试使用不同的整数或浮点值进行 arr_mult 和 arr_div 测试以进行乘法/除法。特别是,比较乘以“2”和乘以“2”。

      【讨论】:

      • 在你的最后一行中,你的意思是dividing by '2' and dividing by '2.0'吗?
      • 这仍然没有回答为什么矢量化版本不显示相同行为的问题......
      • @mtrw 我的意思是“乘法”,尽管除法也可以看到效果。我应该写“2.0”,而不是“2”。
      • @jotasi 我不确定,但考虑到整数除以向量实际上似乎比整数除以浮点数要慢,我会冒险猜测硬件会进行类型转换并在内部进行浮点计算,然后巧妙地转换回来,确保在进行常规整数运算时得到相同的结果。
      【解决方案4】:

      这是第一个操作,通常在“预热”之前需要更长的时间(例如分配的内存、缓存)。

      使用相反的除法和乘法顺序查看相同的效果:

      >>> print_time("arrdiv", timeit.timeit("arrdiv(arr2)", "from __main__ import arrdiv, arr2", number=timeit_iterations))
      >>> print_time("arrmult", timeit.timeit("arrmult(arr2)", "from __main__ import arrmult, arr2", number=timeit_iterations))
      
      arrdiv:  3.2630s
      arrmult:  2.5873s
      

      【讨论】:

      • 实际上,这只是解释了np.float32 数组的微小差异,其中除法已经稍微慢了一点。如果您尝试使用 np.uint8 数组除法仍然快两倍。
      猜你喜欢
      • 2011-05-06
      • 1970-01-01
      • 1970-01-01
      • 2018-05-14
      • 2011-09-22
      • 1970-01-01
      • 2014-08-25
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多