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

如何使用tensordot实现带额外维度的批量矩阵乘法

带批量维度的二维矩阵乘法实现方案

你需要实现的是沿第三维度逐通道独立计算二维矩阵乘积,即对每个通道下标k,单独计算a[:,:,k] @ b[:,:,k],再将所有通道的结果拼接为(2,2,3)的输出。

错误原因说明

  • 原tensordot(a,b,1)默认收缩a的最后一维、b的第一维,当a最后一维为3、b第一维为2时,维度尺寸不匹配触发报错
  • 配置axes=((0,1),(0,1))会完全收缩a、b的前两维,仅保留双方的第三维,输出形状为(3,3),不符合预期

实现方案

方案1:使用np.einsum(最简洁)

直接通过爱因斯坦求和约定指定维度运算规则,无需调整轴顺序:

import numpy as np
# 生成测试输入
a = np.random.rand(2,2,3)
b = np.random.rand(2,2,3)
# 按规则计算乘积
res_einsum = np.einsum('ijk,jlk->ilk', a, b)
print(res_einsum.shape) # 输出 (2, 2, 3)

方案2:调整轴顺序后使用@运算符(可读性更高)

numpy的@运算符原生支持批量维度运算,只需将批量通道维度调整到最前:

# 将第三维移到最前,转为(3,2,2)的批量矩阵格式
a_batch = np.moveaxis(a, -1, 0)
b_batch = np.moveaxis(b, -1, 0)
# 批量矩阵乘法,结果形状为(3,2,2)
res_batch = a_batch @ b_batch
# 将批量维度移回最后,得到(2,2,3)的输出
res_at = np.moveaxis(res_batch, 0, -1)
print(res_at.shape) # 输出 (2, 2, 3)

结果验证

两种推荐方案的输出完全一致:

print(np.allclose(res_einsum, res_at)) # 输出 True

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 20:24:05