【问题标题】:Why dgemm (Cython compiled) is slower than numpy.dot为什么 dgemm(Cython 编译)比 numpy.dot 慢
【发布时间】:2020-04-14 01:31:35
【问题描述】:

长话短说,我在Cython 中构建了一个简单的乘法函数,调用scipy.linalg.cython_blas.dgemm,编译它并针对基准Numpy.dot 运行它。我听说过当我使用静态定义、数组维度预分配、内存视图、关闭检查等技巧时,我将见证 50% 到 100 倍性能提升的神话。但后来我写了自己的 my_dot函数(编译后),它比默认的 Numpy.dot 慢 4 倍。不知道是什么原因,只能猜测:

1) BLAS 库未链接

2) 可能有一些我没有发现的内存开销

3) dot 正在使用一些隐藏的魔法

4) setup.py 写得不好,c 代码没有经过优化编译

5) 我的my_dot 函数没有高效编写

下面是我的代码 sn-p 以及我能想到的所有相关信息可能有助于解决这个难题。如果有人能提供一些关于我做错了什么的见解,或者如何将性能提升到至少与默认的Numpy.dot 持平,我将不胜感激

文件 1:model_cython/multi.pyx。您还需要文件夹中的model_cython/init.py。

#cython: language_level=3 
#cython: boundscheck=False
#cython: nonecheck=False
#cython: wraparound=False
#cython: infertypes=True
#cython: initializedcheck=False
#cython: cdivision=True
#distutils: extra_compile_args = -Wno-unused-function -Wno-unneeded-internal-declaration


from scipy.linalg.cython_blas cimport dgemm
import numpy as np
from numpy cimport ndarray, float64_t
from numpy cimport PyArray_ZEROS
cimport numpy as np
cimport cython

np.import_array()
ctypedef float64_t DOUBLE

def my_dot(double [::1, :] a, double [::1, :] b, int ashape0, int ashape1, 
        int bshape0, int bshape1):
    cdef np.npy_intp cshape[2]
    cshape[0] = <np.npy_intp> ashape0
    cshape[1] = <np.npy_intp> bshape1

    cdef:
        int FORTRAN = 1
        ndarray[DOUBLE, ndim=2] c = PyArray_ZEROS(2, cshape, np.NPY_DOUBLE, FORTRAN)

    cdef double alpha = 1.0
    cdef double beta = 0.0
    dgemm("N", "N", &ashape0, &bshape1, &ashape1, &alpha, &a[0,0], &ashape0, &b[0,0], &bshape0, &beta, &c[0,0], &ashape0)
    return c

文件 2:model_cython/example.py。执行基准测试的脚本

setup_str = """
import numpy as np
from numpy import float64
from multi import my_dot

a = np.ones((2,3), dtype=float64, order='F')
b = np.ones((3,2), dtype=float64, order='F')
print(a.flags)
ashape0, ashape1 = a.shape
bshape0, bshape1 = b.shape
"""
import timeit
print(timeit.timeit(stmt='c=my_dot(a,b, ashape0, ashape1, bshape0, bshape1)', setup=setup_str, number=100000))
print(timeit.timeit(stmt='c=a.dot(b)', setup=setup_str, number=100000))

文件 3:setup.py。编译.so文件

from distutils.core import setup, Extension
from Cython.Build import cythonize
from Cython.Distutils import build_ext
import numpy 
import os
basepath = os.path.dirname(os.path.realpath(__file__))
numpy_path = numpy.get_include()
package_name = 'multi'
setup(
        name='multi',
        cmdclass={'build_ext': build_ext},
        ext_modules=[Extension(package_name, 
            [os.path.join(basepath, 'model_cython', 'multi.pyx')], 
            include_dirs=[numpy_path],
            )],
        )

文件 4:run.sh。执行 setup.py 并移动内容的 Shell 脚本

python3 setup.py build_ext --inplace
path=$(pwd)
rm -r build
mv $path/multi.cpython-37m-darwin.so $path/model_cython/
rm $path/model_cython/multi.c

下面是编译消息的截图:

关于BLAS,我的Numpy 与/usr/local/lib 正确链接,clang -bundle 似乎也在编译中添加-L/usr/local/lib。但也许这还不够?

【问题讨论】:

  • 您正在处理极小的数组。我怀疑大部分运行时将被 Python 调用和验证你的类型占用,而不是你正在做的计算。 (如果你愿意,你也可以在你的函数中获取大小而不是传递它们,尽管我怀疑这会改变速度)
  • 这是我最初的想法,但是numpy.dot 应该做同样的事情。当我将a 和b 传递给dot 时,我没有告诉他们会发生什么。对于my_dot,至少我提供了维度信息。我不知道如何在不进行不公平比较的情况下进一步简化我的功能。而我正在尝试做的实际程序是一个带有 (small n * small m * big S) 的张量循环,所以我试图用小型矩阵优化我的计算。另一件事是这是带有静态类型的 cython,所以我不认为验证类型在这里是个大问题。
  • 但是在每次调用 my_dot 时,Cython 都必须检查类型是否符合您的预期,因为它是从 Python 接收它们并且不能相信它们是正确的。你最好使用cdef 函数(这样你就可以快速从 Cython 调用它)并在 Cython 中的大 S 上循环
  • 1) 使用这个小函数,您基本上可以测量函数调用开销。只有在较大的 3d 阵列上执行此操作时,才有可能实现较大的加速。 2) 对 (2,3)x(3,2) 的 BLAS 调用通常比内联计算慢。 3)想想你额外的编译参数,比如 (-march-native , -fastmath) 看看stackoverflow.com/a/59356461/4045774
  • 对于许多小型数组,您可以比@DavidW 或 einsum 或 @ 运算符的答案快得多。例如。 A=np.random.rand(10_000,3,2)B=np.random.rand(10_000,2,3) 你可以得到dot323(A,B) generated with my gen_dot_nm(3,2,3) -&gt; 18.3 µs 而不是np.einsum("xik,xkj-&gt;xij",A,B)-&gt; 3.99 ms 或A@B -&gt; 3.92 ms 或np.dot(A,B -&gt; 24.8 s。

标签: python c numpy cython matrix-multiplication


【解决方案1】:

Cython 擅长优化循环(这在 Python 中通常很慢),也是调用 C 的便捷方式(这是您想要做的)。但是,从 Python 调用 Cython 函数可能相对较慢 - 特别是因为您指定的所有类型都需要检查一致性。因此,您通常会尝试在一个 Cython 调用后面隐藏大量工作,这样开销就会很小。

您选择了几乎最坏的情况:大量调用背后的一小部分工作。 Cython 或 np.dot 是否会产生更多开销是相当随意的,但无论哪种方式,您测量的都是这个,而不是 np.dot 与 BLAS dgemm。

从您的 cmets 看来,您实际上想要对两个 3D 数组的前两个维度进行点积。因此,一个更有用的测试是尝试重现它。以下是三个版本:

def einsum_mult(a,b):
    # use np.einsum, won't benefit from Cython
    return np.einsum("ijh,jkh->ikh",a,b)

def manual_mult(a,b):
    # multiply one matrix at a time with numpy dot
    # (could probably be optimized a bit with Cython)
    c = np.empty((a.shape[0],b.shape[1],a.shape[2]),
                 dtype=np.float64, order='F')
    for n in range(a.shape[2]):
        c[:,:,n] = a[:,:,n].dot(b[:,:,n])
    return c

def blas_version(double[::1,:,:] a,double[::1,:,:] b):
    # uses dgemm
    cdef double[::1,:,:] c = np.empty((a.shape[0], b.shape[1], a.shape[2]),
                                      dtype=np.float64, order='F')
    cdef double[::1,:] c_part
    cdef int n
    cdef double alpha = 1.0
    cdef double beta = 0.0
    cdef int ashape0 = a.shape[0], ashape1 = a.shape[1], bshape0 = b.shape[0], bshape1 = b.shape[1]

    assert a.shape[2]==b.shape[2]
    assert a.shape[1]==b.shape[0]

    for n in range(a.shape[2]):
        c_part = c[:,:,n]
        dgemm("N", "N", &ashape0, &bshape1, &ashape1, &alpha, &a[0,0,n], &ashape0, 
              &b[0,0,n], &bshape0, &beta, &c_part[0,0], &ashape0)
    return c

使用大小为 (2,3,10000) 和 (3,2,10000) 的数组,重复 100 次,我得到:

manual_mult 1.6531286190001993 s    (i.e. quite bad)
einsum 0.3215398370011826 s         (pretty good)
blas_version 0.15762194800481666 s  (best, pretty close to the "myth" performance gain you mention)

如果您充分利用 Cython 并将循环保留在已编译代码中,则 BLAS 版本会很快。 (我没有花精力优化这个,所以如果你尝试过,你可能会打败它,但它只是为了说明这一点)

【讨论】:

  • 感谢您的出色回答!是的,这正是我最终想要做的。我首先建立一个小的my_fun 的原因是我以前从未做过cython,所以我只是想确保我不会在计算量最大的操作上做愚蠢的事情。对组件进行模块化测试,一次检查一个。但我想我毕竟做了愚蠢的事情,哈哈。再次感谢您的回答。
  • 您介意与我分享您的setup.py 和.pyx 代码的序言(例如那些#cython: xxxx)吗?我听说代码的编译方式也会影响性能。只想向有经验的人学习。
  • 你会失望的。我只是在命令行上运行:cythonize -i filename.pyx(就地构建)。我怀疑这是错误的 - 我应该与 BLAS 链接,但因为我先导入 Numpy,所以我侥幸逃脱了。它只是作为一个快速演示编写的......
  • 我绝对没有使用#cython 指令。我认为关闭boundschecking 和wraparound 可能会有所帮助。我认为任何其他指令都不会产生太大/任何差异。我的观点是,您应该使编译指令尽可能本地化,尽可能少,并且仅在您需要它们时才应用(即仅在您知道它可以执行某些操作并且已放入我的 @987654337 之类的东西时才应用 boundscheck @ 声明以确保它不会导致崩溃)
猜你喜欢
  • 1970-01-01
  • 2016-10-09
  • 1970-01-01
  • 2014-08-10
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2011-08-25
  • 1970-01-01
相关资源
最近更新 更多