You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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性能
101.7 µs ± 3.79 ns/循环(10次测试,每次10000000循环)2.0 µs ± 43.28 ns/循环(10次测试,每次10000000循环)
5083.5 µs ± 398.63 ns/循环(10次测试,每次100000循环)84.7 µs ± 666.96 ns/循环(10次测试,每次100000循环)
100123.8 µs ± 3.56 µs/循环(10次测试,每次100000循环)124.1 µs ± 1.43 µs/循环(10次测试,每次100000循环)
5004.8 ms ± 58.56 µs/循环(10次测试,每次10000循环)5.1 ms ± 6.00 µs/循环(10次测试,每次1000循环)
100031.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性能
101.7 µs ± 3.88 ns/循环(10次测试,每次10000000循环)1.1 µs ± 34.48 ns/循环(10次测试,每次10000000循环)
5089.1 µs ± 2.44 µs/循环(10次测试,每次1000000循环)89.3 µs ± 208.79 ns/循环(10次测试,每次1000000循环)
100135.5 µs ± 105.27 ns/循环(10次测试,每次100000循环)127.9 µs ± 37.27 ns/循环(10次测试,每次100000循环)
5004.8 ms ± 3.67 µs/循环(10次测试,每次10000循环)5.2 ms ± 4.33 µs/循环(10次测试,每次10000循环)
100031.8 ms ± 55.75 µs/循环(10次测试,每次1000循环)33.4 ms ± 62.54 µs/循环(10次测试,每次1000循环)

内容的提问来源于stack exchange,提问作者Maryks

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 18:54:52