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

如何利用Broadcasting消除3D矩阵批量乘法中的循环?

利用NumPy广播消除矩阵乘法循环

问题场景

初始单个矩阵乘法代码:

import numpy as np
np.random.seed(2)

result = SC @ x

其中SC是nn x nn维度矩阵,x是nn x ns维度矩阵。

现在我们有一个3D张量SCs,维度为ns x nn x nn,原始循环实现如下:

ns = 4
nn = 2

SCs = np.random.rand(ns, nn, nn)
x = np.random.rand(nn, ns)

def matmul3d(a, b):
    ns, nn, nn = a.shape
    assert(b.shape == (nn, ns))
    
    results = np.zeros((nn, ns))
    for i in range(ns):
        results[:, i] = a[i, :, :] @ b[:, i]
    return results

该函数输出:

array([[0.385428  , 0.22932766, 0.36791082, 0.06029485],
       [0.68934311, 0.14157493, 0.75236553, 0.09049892]])

直接执行a @ b会得到ns x nn x ns的张量,其中对角线位置的元素(对应循环中i索引的结果)就是我们需要的目标输出:

full_result = SCs @ x

输出:

array([[[0.385428  , 0.21717737, 0.38019609, 0.0372277 ],
        [0.68934311, 0.30008412, 0.65169432, 0.0858002 ]],

       [[0.52588409, 0.22932766, 0.4972909 , 0.06536792],
        [0.48764911, 0.14157493, 0.43837138, 0.07607813]],

       [[0.39071113, 0.1655206 , 0.36791082, 0.04962322],
        [0.79777992, 0.34153306, 0.75236553, 0.10054907]],

       [[0.37441129, 0.10004409, 0.33380446, 0.06029485],
        [0.5542946 , 0.14242876, 0.4923592 , 0.09049892]]])

广播消除循环的实现方法

完全可以利用NumPy的广播机制替代循环,以下是两种高效实现方式:

方法1:形状适配+批量矩阵乘法

通过调整x的形状适配广播规则,再提取结果:

# 将x调整为(ns, nn, 1),适配SCs的维度广播
x_reshaped = x.T[:, :, np.newaxis]
# 批量矩阵乘法:(ns, nn, nn) @ (ns, nn, 1) → (ns, nn, 1)
temp = SCs @ x_reshaped
# 挤压维度并转置得到目标形状(nn, ns)
result = temp.squeeze().T

方法2:Einstein求和(更简洁)

用np.einsum直接描述运算逻辑,无需手动调整形状:

result = np.einsum('ijk,ik->ji', SCs, x)

表达式解释:

  • ijk对应SCs的维度:ns(i) x nn(j) x nn(k)
  • ik对应x的维度:nn(k) x ns(i)
  • ->ji指定输出维度为nn(j) x ns(i),完全匹配目标结果的形状。

两种方法的输出都和原循环函数完全一致,且依赖NumPy底层优化,运算速度远快于Python原生循环,尤其在ns、nn数值较大时优势明显。

内容的提问来源于stack exchange,提问作者Abolfazl

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 02:05:37