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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 14:27:16