【发布时间】: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) -> 18.3 µs而不是np.einsum("xik,xkj->xij",A,B)-> 3.99 ms或A@B -> 3.92 ms或np.dot(A,B -> 24.8 s。
标签: python c numpy cython matrix-multiplication