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

如何优化Numba矩阵乘法性能?性能落后Numpy问题排查

你的Numba矩阵乘法为啥比Numpy慢这么多?

别慌,我帮你拆解下问题——你的实现主要踩了内存访问模式和并行策略的坑,这两个是Numba性能优化的核心点,下面给你详细分析和优化方案:

核心问题在哪里?

1. 循环顺序完全搞反了,缓存直接炸了

Numpy数组是行优先存储的,简单说就是同一行的元素在内存里是挨着的,访问起来快。你现在的循环顺序是i → j → k:

  • 访问A[i,k]的时候,k逐个加1,是顺着A的第i行读,缓存能命中,没问题;
  • 但访问B[k,j]的时候,j固定,k逐个加1,相当于读B的第j列,这元素在内存里是跳着放的,缓存根本存不住,每次都要从内存里读,速度直接慢10倍都不止,这是性能差的头号原因。

2. prange用错地方了,越帮越忙

你在j循环上用了prange,这就像给蚂蚁开卡车——j循环的粒度太小了,线程创建、切换的开销比计算本身还大,完全发挥不了并行的优势。而且这种细粒度并行还可能带来缓存竞争,进一步拖慢速度。

3. 还有个潜在的维度bug

你的k循环范围是A.shape[0],但矩阵乘法里k应该遍历A的列数(也就是B的行数),也就是A.shape[1]。你现在用方阵测试没报错,换成非方阵直接算错,得先把这个坑填上。

优化后的代码,性能直接追平Numpy

下面是修正后的版本,把所有坑都填上了:

import numpy as np
import timeit
from numba import jit, float64, prange

@jit(float64[:,:](float64[:,:], float64[:,:]), parallel=True, nopython=True, fastmath=True)
def matmul(A, B):
    # 先做维度校验,避免踩坑
    assert A.shape[1] == B.shape[0], "矩阵维度不匹配,没法相乘"
    C = np.zeros((A.shape[0], B.shape[1]), dtype=float64)
    # 关键:把循环顺序改成i → k → j,保证内存连续访问
    for i in prange(A.shape[0]):
        for k in range(A.shape[1]):
            # 提前把A[i,k]取出来,减少重复索引访问的开销
            a_val = A[i, k]
            for j in range(B.shape[1]):
                C[i, j] += a_val * B[k, j]
    return C

if __name__ == '__main__':
    m_size = 1000
    num_loops = 10
    A = np.random.rand(m_size, m_size)
    B = np.random.rand(m_size, m_size)
    
    # 重要:预热Numba函数!第一次调用会编译,别把编译时间算进测试里
    matmul(A[:2,:2], B[:2,:2])
    
    # Numpy基准测试
    start = timeit.default_timer()
    for _ in range(num_loops):
        A.dot(B)
    print(f"Numpy 耗时: {timeit.default_timer() - start:.6f} 秒")
    
    # Numba测试
    start = timeit.default_timer()
    for _ in range(num_loops):
        matmul(A, B)
    print(f"Numba 耗时: {timeit.default_timer() - start:.6f} 秒")

额外的优化小技巧

  • 预热不能忘:Numba第一次运行函数会编译,这个时间很长,一定要先跑个小矩阵预热,不然测试结果不准。
  • 开fastmath:加上fastmath=True,允许Numba用一些近似运算(忽略NaN/Inf的特殊处理),性能能再提一截,如果你不需要严格的IEEE标准的话。
  • 别硬写类型注解:如果不确定矩阵类型,直接用@jit(parallel=True, nopython=True, fastmath=True)让Numba自己推断,代码更灵活。
  • 大矩阵可以试试分块:如果矩阵特别大(比如2000x2000以上),可以把矩阵切成小块计算,进一步提升缓存利用率,这就是Numpy底层BLAS库用的技巧。

改完之后你再跑,Numba的性能应该和Numpy的dot差不多,甚至在某些场景下能反超~

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:27:29