【问题标题】:Numba is not enhancing the performanceNumba 没有提高性能
【发布时间】:2022-01-24 02:01:06
【问题描述】:

我正在测试一些采用 numpy 数组的函数的 numba 性能,并进行比较:

import numpy as np
from numba import jit, vectorize, float64
import time
from numba.core.errors import NumbaWarning
import warnings

warnings.simplefilter('ignore', category=NumbaWarning)

@jit(nopython=True, boundscheck=False) # Set "nopython" mode for best performance, equivalent to @njit
def go_fast(a):     # Function is compiled to machine code when called the first time
    trace = 0.0
    for i in range(a.shape[0]):   # Numba likes loops
        trace += np.tanh(a[i, i]) # Numba likes NumPy functions
    return a + trace              # Numba likes NumPy broadcasting
   
class Main(object):
    def __init__(self) -> None:
        super().__init__()
        self.mat     = np.arange(100000000, dtype=np.float64).reshape(10000, 10000)

    def my_run(self):
        st = time.time()
        trace = 0.0
        for i in range(self.mat.shape[0]):   
            trace += np.tanh(self.mat[i, i]) 
        res = self.mat + trace
        print('Python Diration: ', time.time() - st)
        return res                           
    
    def jit_run(self):
        st = time.time()
        res = go_fast(self.mat)
        print('Jit Diration: ', time.time() - st)
        return res
        
obj = Main()
x1 = obj.my_run()
x2 = obj.jit_run()

输出是:

Python Diration:  0.2164750099182129
Jit Diration:  0.5367801189422607

如何获得此示例的增强版本?

【问题讨论】:

  • 排除 Numba 的编译时间(即忽略 JIT 函数的首次运行)时,我无法在我的机器上重现问题:两者都需要大约 0.1 秒。
  • 这就是答案。在我的机器上测试,我得到Python duration: 0.23 然后Jit duration: 0.79 0.20 0.20 0.20 0.20 ...
  • 你只计时第一次运行吗?
  • 是的,我正在计时第一次运行。有什么方法可以在不运行的情况下初始化 jit 函数?因为我将在正常工作中运行一次该功能

标签: python numpy numba jit


【解决方案1】:

Numba 实现较慢的执行时间是由于编译时间,因为 Numba 在函数使用时编译函数(除非参数类型发生变化,仅在第一次编译)。它这样做是因为它在调用函数之前无法知道参数的类型。希望您可以为 Numba 指定参数类型,以便它可以直接编译函数(在执行装饰器函数时)。这是生成的代码:

@njit('float64[:,:](float64[:,:])')
def go_fast(a):
    trace = 0.0
    for i in range(a.shape[0]):
        trace += np.tanh(a[i, i])
    return a + trace

请注意,njitjit+nopython=True 的快捷方式,并且boundscheck 默认已设置为False(请参阅doc)。

在我的机器上,这导致 Numpy 和 Numba 的执行时间相同。实际上,执行时间不受tanh 函数计算的限制。它以表达式a + trace 为界(适用于 Numba 和 Numpy)。预计执行时间相同,因为两者都以相同的方式实现:它们创建一个临时的新数组来执行加法。由于page faults 和RAM 的使用(a 完全从 RAM 中读取并且临时数组完全存储在 RAM 中),创建一个新的临时数组的成本很高。如果您想要更快的计算,那么您需要执行 in-place 操作(这可以防止在 x86 平台上出现页面错误和昂贵的cache-line write allocations)。

【讨论】:

  • 感谢您的解释。我有三个问题,1)您如何阅读此表达式 'float64[:,:](float64[:,:])'(这是什么意思),2)如果我有第二个输入(类型为 numba.type。 pyobject),如何将其添加到此表达式中,以及 3)如何将输出的类型添加到此表达式中(假设它是一个相同类型的 float64 数组)?
  • 1) 语法是ReturnType(Arg1_Type, Arg2_Type, ...)。一维数组的类型是ItemType[:]。对于 2D 数组,它是 ItemType[:, :] 等等。这里项目的类型是 64 位浮点数(因此是 float64)。请参阅doc 了解更多信息。 2)您不能在 Numba nopython jitted 代码中使用 CPython 对象,因为不使用慢速 CPython 动态类型是 Numba 快速的原因(它使用静态类型并且没有 GC)。否则你可以使用pyobjectnopython=False (但我这不是很有用)。
  • 对于问题3,我想答案1已经给出了答案:括号前的第一个参数是返回类型。
  • 我对 Arg1(numpy array) 和 Arg2(pyobject) 使用了如下语法:@jit(('float64[:,:](float64[:,:])', 'pyobject '), nopython=False).....-> 给我这个错误 [TypeError: invalid type in signature: expected a type instance, got 'float64[:,:](float64[:,:])']
  • 我不确定您到底想要什么,但我认为您需要:@njit('float64[:,:](float64[:,:], pyobject)')
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2021-09-06
  • 1970-01-01
  • 2014-06-26
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多