Cython实现矩阵乘法报错,_memoryviewslice不支持@运算符如何解决?
在Cython中实现矩阵乘法解决@运算符错误
你的问题出在Cython的memoryview类型不支持@矩阵乘法运算符,下面提供几种可行的解决方式:
方法1:将Memoryview转换为NumPy数组,复用原生矩阵乘法
这是最简单的解决方案,直接把memoryview转回NumPy数组,就能使用NumPy优化过的矩阵乘法:
import numpy as np def cell(float[:, ::1] a, float[:, ::1] b): # 将memoryview转换为numpy数组 np_a = np.asarray(a) np_b = np.asarray(b) c = np_a @ np_b return c
这种方式几乎不需要额外代码,而且NumPy的矩阵乘法底层已经做了优化,对于大多数场景速度足够。
方法2:手动实现C级矩阵乘法(追求极致性能)
如果需要进一步提速,可以手动编写Cython循环直接操作内存,同时关闭边界检查来减少开销:
import numpy as np cimport numpy as np from cython cimport boundscheck, wraparound @boundscheck(False) # 关闭边界检查,提升速度 @wraparound(False) # 关闭负索引支持,提升速度 def cell(np.float32_t[:, ::1] a, np.float32_t[:, ::1] b): cdef int m = a.shape[0] cdef int k = a.shape[1] cdef int n = b.shape[1] # 检查矩阵维度是否匹配 if k != b.shape[0]: raise ValueError("矩阵维度不匹配:a的列数不等于b的行数") # 初始化结果矩阵 cdef np.ndarray[np.float32_t, ndim=2] c = np.zeros((m, n), dtype=np.float32) cdef np.float32_t[:, ::1] c_view = c cdef int i, j, l cdef float temp # 三重循环实现矩阵乘法 for i in range(m): for j in range(n): temp = 0.0 for l in range(k): temp += a[i, l] * b[l, j] c_view[i, j] = temp return c
这种方式避免了Python层的开销,适合小矩阵或对延迟敏感的场景,但代码复杂度更高。
方法3:调用BLAS库的优化函数(最优性能)
如果你的系统安装了BLAS库(比如OpenBLAS),可以直接调用其高度优化的矩阵乘法函数,这是大矩阵场景下速度最快的方案:
import numpy as np cimport numpy as np from cython cimport boundscheck, wraparound cimport openblas # 根据你的BLAS库选择,比如cimport blas @boundscheck(False) @wraparound(False) def cell(np.float32_t[:, ::1] a, np.float32_t[:, ::1] b): cdef int m = a.shape[0] cdef int k = a.shape[1] cdef int n = b.shape[1] if k != b.shape[0]: raise ValueError("矩阵维度不匹配") cdef np.ndarray[np.float32_t, ndim=2] c = np.zeros((m, n), dtype=np.float32) # 调用OpenBLAS的单精度矩阵乘法函数sgemm # 参数说明:转置标识, 转置标识, m, n, k, alpha, A指针, A的列数, B指针, B的列数, beta, C指针, C的列数 openblas.sgemm( b'N', b'N', # A和B都不转置 m, n, k, 1.0, &a[0,0], k, &b[0,0], n, 0.0, &c[0,0], n ) return c
使用前需要确保Cython能正确链接到BLAS库,编译时可能需要添加对应的编译参数(比如-lopenblas)。
内容的提问来源于stack exchange,提问作者Gustavo Caetano
相关产品推荐
相关产品推荐

