如何让Cython+cython_blas矩阵乘法性能媲美numpy.dot?
复数矩阵乘法的Cython优化
需求与基准实现
我需要对两个复数矩阵执行乘法运算,基础代码如下:
import numpy as np n_dim = ... # 指定矩阵维度 a = np.random.randn(n_dim, n_dim) + 1j * np.random.randn(n_dim, n_dim) b = np.random.randn(n_dim, n_dim) + 1j * np.random.randn(n_dim, n_dim)
我以np.dot(a, b)作为基准实现,推测它底层调用了BLAS的zgemm函数,因此希望编写一个直接调用zgemm的Cython函数,性能至少不弱于np.dot。
初始Cython实现
# cython: language_level=3 # cython: infer_types=False import numpy as np cimport numpy as np cimport scipy.linalg.cython_blas as blas cimport cython @cython.boundscheck(False) @cython.wraparound(False) def func(double complex[:, ::1] a, double complex[:, ::1] b): cdef: int m, n, k, lda, ldb, ldc double complex[:, ::1] cv double complex alpha, beta m = a.shape[0] n = b.shape[1] k = a.shape[1] lda = a.shape[0] ldb = b.shape[0] ldc = a.shape[0] alpha = 1 + 0j beta = 0 + 0j c = np.empty((n, m), dtype=np.complex128) cv = c # zgemm计算逻辑: C = alpha * op(A) * op(B) + beta * C # 参数顺序: TRANSA, TRANSB, M, N, K, ALPHA, A, LDA, B, LDB, BETA, C, LDC blas.zgemm('N', 'N', &n, &m, &k, &alpha, &b[0, 0], &ldb, &a[0, 0], &lda, &beta, &cv[0, 0], &ldc) return c
注:由于输入数组是C顺序,而zgemm更适配Fortran顺序,我通过计算(b.T @ a.T).T来避免数组复制,而非直接计算a @ b。
初始性能测试
| 矩阵维度 | np.dot性能 | func性能 |
|---|---|---|
| 10 | 1.7 µs ± 3.79 ns/循环(10次测试,每次10000000循环) | 2.0 µs ± 43.28 ns/循环(10次测试,每次10000000循环) |
| 50 | 83.5 µs ± 398.63 ns/循环(10次测试,每次100000循环) | 84.7 µs ± 666.96 ns/循环(10次测试,每次100000循环) |
| 100 | 123.8 µs ± 3.56 µs/循环(10次测试,每次100000循环) | 124.1 µs ± 1.43 µs/循环(10次测试,每次100000循环) |
| 500 | 4.8 ms ± 58.56 µs/循环(10次测试,每次10000循环) | 5.1 ms ± 6.00 µs/循环(10次测试,每次1000循环) |
| 1000 | 31.3 ms ± 48.87 µs/循环(10次测试,每次1000循环) | 32.9 ms ± 153.06 µs/循环(10次测试,每次1000循环) |
两者性能接近,但我预期func会因为检查步骤更少而始终不弱于np.dot,实际却并非如此。我想知道如何进一步优化func的速度,比如是否存在更优的输出数组分配方式?
优化后的Cython实现
将np.empty替换为NumPy C API的np.PyArray_EMPTY,并用np.ndarray[np.complex128_t, ndim=2]替代double complex[:, ::1]后,性能有所提升:
# cython: language_level=3 # cython: infer_types=False import numpy as np cimport numpy as np cimport scipy.linalg.cython_blas as blas cimport cython np.import_array() @cython.boundscheck(False) @cython.wraparound(False) def func1(np.ndarray[np.complex128_t, ndim=2] a, np.ndarray[np.complex128_t, ndim=2] b): cdef: int k, m, n double complex alpha, beta m = a.shape[0] k = a.shape[1] n = b.shape[1] alpha = 1 + 0j beta = 0 + 0j cdef: np.npy_intp *dims = [m, n] np.ndarray[np.complex128_t, ndim=2] c = np.PyArray_EMPTY(2, dims, np.NPY_COMPLEX128, 0) blas.zgemm('N', 'N', &n, &m, &k, &alpha, &b[0, 0], &n, &a[0, 0], &k, &beta, &c[0, 0], &n) return c
优化后性能测试
| 矩阵维度 | np.dot性能 | func1性能 |
|---|---|---|
| 10 | 1.7 µs ± 3.88 ns/循环(10次测试,每次10000000循环) | 1.1 µs ± 34.48 ns/循环(10次测试,每次10000000循环) |
| 50 | 89.1 µs ± 2.44 µs/循环(10次测试,每次1000000循环) | 89.3 µs ± 208.79 ns/循环(10次测试,每次1000000循环) |
| 100 | 135.5 µs ± 105.27 ns/循环(10次测试,每次100000循环) | 127.9 µs ± 37.27 ns/循环(10次测试,每次100000循环) |
| 500 | 4.8 ms ± 3.67 µs/循环(10次测试,每次10000循环) | 5.2 ms ± 4.33 µs/循环(10次测试,每次10000循环) |
| 1000 | 31.8 ms ± 55.75 µs/循环(10次测试,每次1000循环) | 33.4 ms ± 62.54 µs/循环(10次测试,每次1000循环) |
内容的提问来源于stack exchange,提问作者Maryks
相关产品推荐
相关产品推荐

