显式2x2矩阵乘法为何比numpy.matmul快一个数量级?
为什么显式实现2x2批量矩阵乘法比Numpy内置方法快10倍?
我需要对两个形状为(n,2,2)的Numpy数组执行批量逐矩阵乘法(即第i个结果矩阵是输入数组中第i对矩阵的乘积)。性能分析显示矩阵乘法是瓶颈,于是我手动实现了2x2矩阵的乘法逻辑,测试后发现它比@、matmul甚至einsum快近10倍,但结果完全一致。
测试代码如下:
import timeit import numpy as np def explicit_2x2_matrices_multiplication( mats_a: np.ndarray, mats_b: np.ndarray ) -> np.ndarray: matrices_multiplied = np.empty_like(mats_b) for i in range(2): for j in range(2): matrices_multiplied[:, i, j] = ( mats_a[:, i, 0] * mats_b[:, 0, j] + mats_a[:, i, 1] * mats_b[:, 1, j] ) return matrices_multiplied matrices_a = np.random.random((1000, 2, 2)) matrices_b = np.random.random((1000, 2, 2)) assert np.allclose( # 验证显式实现的正确性 matrices_a @ matrices_b, explicit_2x2_matrices_multiplication(matrices_a, matrices_b), ) print( # 1.1814142999992328 秒 timeit.timeit(lambda: matrices_a @ matrices_b, number=10000) ) print( # 1.1954495010013488 秒 timeit.timeit(lambda: np.matmul(matrices_a, matrices_b), number=10000) ) print( # 2.2304022700009227 秒 timeit.timeit(lambda: np.einsum('lij,ljk->lik', matrices_a, matrices_b), number=10000) ) print( # 0.19581600800120214 秒 timeit.timeit( lambda: explicit_2x2_matrices_multiplication(matrices_a, matrices_b), number=10000, ) )
性能差距的核心原因:
- 通用实现的开销占比过高:
matmul/@是面向任意维度、任意大小矩阵的通用实现,包含大量参数校验、维度适配、计算路径选择(比如是否调用BLAS库)的逻辑。这些额外操作对于极小的2x2矩阵来说,开销占比远高于实际计算成本,而显式实现完全跳过了这些通用逻辑,直接针对固定结构计算。 - 向量化效率与内存访问优化:显式实现中的所有计算都是针对一维数组(如
mats_a[:,i,0]是形状为(n,)的数组)的逐元素乘加,Numpy能直接将这些操作映射到CPU的SIMD指令,且内存访问模式连续友好。而通用矩阵乘法对2x2矩阵无法充分发挥BLAS的优化优势(BLAS更适配大矩阵),甚至BLAS的初始化开销会超过计算本身。 einsum的解析成本:einsum需要先解析下标字符串、构建计算图,这一步的开销对于小矩阵批量计算来说非常显著,因此它的速度最慢。- Python循环的“假开销”:代码中的两层Python循环仅执行4次,循环次数极少,几乎没有Python级别的循环开销(Python循环慢是针对大量循环的场景)。反而固定的循环次数避免了通用实现中的分支判断逻辑。
内容的提问来源于stack exchange,提问作者TheFamousRat
相关产品推荐
相关产品推荐

