如何对NumPy多组连续矩阵乘积实现向量化运算?
向量化实现批量矩阵连乘方案
针对你提到的形状为(M, L, 2, 2)的NumPy数组,要批量对每个M对应的L个2x2矩阵做连乘,完全可以用向量化方式实现,不用手动遍历M和L,具体方案如下:
- 核心思路:利用
functools.reduce结合NumPy的矩阵乘法@,通过调整数组轴的顺序,让批量矩阵连乘以向量化方式执行。
代码示例
import numpy as np from functools import reduce # 生成测试数据(示例中M=1000,L=3) M = 1000 L = 3 ar = np.random.rand(M, L, 2, 2) # 执行向量化连乘 # 先把L所在的轴(轴1)移到最前面,数组形状变为(L, M, 2, 2) # 再用reduce依次对L组矩阵做批量乘法 result = reduce(lambda x, y: x @ y, np.moveaxis(ar, 1, 0))
原理说明
np.moveaxis(ar, 1, 0)将原数组的形状从(M, L, 2, 2)转为(L, M, 2, 2),此时每个元素是形状为(M, 2, 2)的数组,对应所有M样本的第k个矩阵。reduce会依次将这些(M, 2, 2)的数组做矩阵乘法@,这个乘法是批量执行的——对每个M样本,各自对应的矩阵进行相乘,最终得到形状为(M, 2, 2)的结果数组。
验证正确性
你可以手动计算单个样本的结果来验证向量化方案的准确性:
# 取第一个样本手动连乘 manual_res = ar[0, 0] @ ar[0, 1] @ ar[0, 2] # 对比向量化结果与手动计算结果 assert np.allclose(result[0], manual_res)
补充:若L为固定小值
比如L=4,也可以直接展开写法,效率和reduce方案差不多:
result = ar[:,0] @ ar[:,1] @ ar[:,2] @ ar[:,3]
这种方式完全避免了Python层面的循环,依赖NumPy内部的高效C实现,适合M较大的场景,在L≤5的情况下性能最优。
内容的提问来源于stack exchange,提问作者VaNa
相关产品推荐
相关产品推荐

