Python中如何不使用for循环对4D NumPy数组逐行执行矩阵乘法
NumPy 4D数组批量块矩阵乘法实现
你的输入数组k5是形状为(m, n, d, d)的4维数组,其中m是块行数、n是块列数、d是单个小矩阵的边长,需求是对每个块行,将该行的n个小矩阵按顺序做矩阵乘法,最终输出形状为(m, d, d)的结果数组。
测试用例快速实现
针对你给出的n=2(每行2个小矩阵)的测试场景,直接用NumPy原生支持批量维度的矩阵乘法运算符@即可,完全满足无显式for循环、不调用np.linalg.multi_dot的要求:
import numpy as np k1=np.array([[1,2],[3,4]]) k2=np.array([[5,6],[7,8]]) k3=np.array([[9,10],[11,12]]) k4=np.array([[13,14],[15,16]]) k5=np.array([[k1,k2],[k3,k4]]) result = k5[:, 0] @ k5[:, 1]
运行后得到的result和你给出的预期输出完全一致:
[[[ 19 22] [ 43 50]] [[267 286] [323 346]]]
实现原理:k5[:,0]取出所有块行的第一个小矩阵,形状为(m, d, d),k5[:,1]取出所有块行的第二个小矩阵,形状同为(m, d, d),@运算符会自动把前导的m维度识别为批量维度,一次性并行完成所有m组矩阵乘法,没有Python层循环开销。
通用场景(任意n个块连乘)实现
如果实际场景中每个块行有n个可连乘的矩阵(n≥2),可以使用np.einsum实现高效批量运算,写法非常直观:
- n=2场景(和上面
@运算符等价):
result = np.einsum('...ij,...jk->...ik', k5[:,0], k5[:,1])
- n=3场景(每行3个矩阵连乘):
result = np.einsum('...ij,...jk,...kl->...il', k5[:,0], k5[:,1], k5[:,2])
- 任意n的场景只需要按矩阵乘法的下标传递规则,扩展einsum的下标字符串、传入对应位置的块矩阵即可。
方案优势
- 所有运算都在C层执行,无Python层面的for循环,运行效率极高
- 不需要依赖
np.linalg.multi_dot接口 - 对小矩阵的尺寸没有强制2×2的限制,只要相邻矩阵满足矩阵乘法的维度匹配要求即可使用
内容的提问来源于stack exchange,提问作者FlamingosAreSad
相关产品推荐
相关产品推荐

