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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 00:25:26