如何优化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
相关产品推荐
相关产品推荐

