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

Numpy中泛化实现(n,3,3)与(n,3)数组逐批次矩阵乘法

批量矩阵-向量乘法泛化实现

你的原始代码通过逐索引循环,对每个样本下的3×3矩阵和3维向量做乘积,要泛化到任意n的批量输入,不需要显式写循环,有两种非常简洁的实现方式:

方案1:使用einsum(你提到的方案)

einsum的下标规则非常贴合这个计算逻辑,直接定义输入输出的维度对应关系即可:

  • 输入a的维度标记为nij:n是批量维度,i、j对应3×3矩阵的两个维度
  • 输入b的维度标记为nj:n是批量维度,j对应3维向量的维度
  • 输出标记为ni:保留批量维度n和结果向量的维度i,对重复出现的j维度自动求和
import numpy as np

# 测试数据和你原始代码一致
a = np.arange(18).reshape(2,3,3)
b = np.arange(6).reshape(2,3)

c = np.einsum('nij,nj->ni', a, b)

运行结果和你原始逐行赋值的代码完全一致,且支持任意正整数n的输入,不需要修改代码逻辑。

方案2:使用原生广播矩阵乘法

如果不想记einsum的下标语法,直接用numpy原生的@运算符也可以实现:numpy的矩阵乘法默认对前置的批量维度做广播,只需要把形状为(n,3)的b扩展一个尾部维度变成(n,3,1),和形状为(n,3,3)的a相乘后会得到形状为(n,3,1)的结果,再压缩掉最后一个长度为1的维度即可得到(n,3)的目标输出:

c = (a @ b[..., np.newaxis]).squeeze(axis=-1)

两种实现都是向量化运算,执行效率远高于显式Python循环,结果完全等价于你原始的逐索引赋值逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 00:33:35