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

Numpy张量中逐矩阵点积的优雅实现方法问询

优雅实现Numpy张量的逐矩阵点积操作

嘿,这个问题一点都不简单!其实Numpy里有好几种非常优雅的方式来实现你要的逐batch矩阵点积操作,我给你列几个最常用的:

方法1:使用np.matmul(推荐)

np.matmul是Numpy专门为矩阵乘法设计的函数,对于高维张量,它会自动将前面的维度视为batch维度,只对最后两个维度执行矩阵乘法,完美匹配你的需求:

import numpy as np
M = np.array(range(12)).reshape(2,2,3)
N = np.array(range(12)).reshape(2,3,2)

# 逐batch矩阵点积
result = np.matmul(M, N)

# 验证结果
print("result[0]等于np.dot(M[0], N[0]):", np.array_equal(result[0], np.dot(M[0], N[0])))
print("result[1]等于np.dot(M[1], N[1]):", np.array_equal(result[1], np.dot(M[1], N[1])))

输出结果的形状是(2,2,2),其中每个result[i]就是你要的np.dot(M[i], N[i])的结果。

方法2:使用@运算符(更简洁)

@是np.matmul的中缀语法糖,写法更简洁直观,功能完全一致:

result = M @ N

这种写法在代码里更清爽,适合日常使用。

方法3:使用np.einsum(灵活性拉满)

如果你需要更精细地控制维度运算,np.einsum是绝佳选择,它通过爱因斯坦求和约定来定义运算,可读性和灵活性都很强:

result = np.einsum('bij,bjk->bik', M, N)

这里的字符串'bij,bjk->bik'可以理解为:

  • b:batch维度(对应你的第一个维度,共2个batch)
  • i,j:M的矩阵维度(2行3列)
  • j,k:N的矩阵维度(3行2列)
  • 箭头后的bik表示输出每个batch下的2行2列矩阵,正好是矩阵乘法的结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:05:17