NumPy中如何对shape为(5,5,3)的三维矩阵正确执行乘法运算
NumPy三维张量乘法tensordot的axes参数配置
首先明确:高维张量不存在通用的"矩阵乘法"定义,tensordot的axes参数作用是指定两个张量上需要做乘加收缩(也就是对应元素相乘再求和)的维度,只要两个维度长度一致就可以配对收缩。
先修正你定义数组的小问题:原代码第四行有个元素写的是元组(0,0,0.25),NumPy虽然会自动转成ndarray不影响运行,但统一写列表形式更稳妥。
axes参数支持两种传参格式:
- 传单个整数
n:自动取第一个张量的最后n个维度、第二个张量的最前n个维度,按顺序配对做收缩 - 传长度为2的元组
(a_axes, b_axes):a_axes是第一个张量要收缩的维度序号列表,b_axes是第二个张量要收缩的维度序号列表,两个列表对应位置的维度长度必须相等
针对你shape为(5,5,3)的Z数组,下面是三种最常见需求对应的参数配置:
场景1:对最后一维长度为3的向量做内积
如果你需要对每个位置存储的3维向量两两做点积(对应元素相乘求和),就把两个张量的最后一维(序号为2)配对收缩,参数设置为axes=([2], [2]):
import numpy as np Z = np.array([ [[0,0,0.25],[0.25,0.5,0.75],[0,0,0.25],[0.75,1.0,1.0],[0.75,1.0,1.0]], [[0,0,0.25],[0,0,0.25],[0.5,0.75,1.0],[0,0,0.25],[0,0,0.25]], [[0,0,0.25],[0,0,0.25],[0,0,0.25],[0,0.25,0.5],[0,0,0.25]], [[0,0,0.25],[0.25,0.5,0.75],[0,0,0.25],[0,0,0.25],[0,0,0.25]], [[0,0,0.25],[0,0,0.25],[0,0,0.25],[0,0,0.25],[0,0,0.25]] ]) res = np.tensordot(Z, Z, axes=([2], [2])) print(res.shape) # 输出 (5, 5, 5, 5)
得到的结果是4维数组,每个位置res[i,j,k,l]的值就是Z[i,j,:]和Z[k,l,:]两个3维向量的内积。
场景2:按最后一维做批量二维矩阵乘法
如果你想把最后一个长度为3的维度当成独立的通道/批次维度,每个通道单独做两个5×5二维矩阵的标准矩阵乘法,不需要用tensordot,直接用np.matmul(也就是@运算符)即可:
res = Z @ Z print(res.shape) # 输出 (5, 5, 3)
你之前用
np.dot(Z,Z)得到错误结果,是因为np.dot处理高于2维的数组时,会把两个数组的最后两个维度当成矩阵维度做乘,最终输出shape为(5,5,3,5,3),不符合批量矩阵乘的预期。
场景3:按5长度维度做矩阵乘收缩
如果你需要把两个数组中长度为5的维度配对做矩阵乘(比如第一个数组的第二维、第二个数组的第一维配对收缩),参数设置为axes=([1], [0]),最终输出shape为(5,3,5,3)。
配置参数时只要记住:你想通过乘加求和消掉哪两个等长维度,就把这两个维度的序号分别放进两个列表传给axes,所有没被指定收缩的维度,会按原顺序保留在输出结果里。
内容的提问来源于stack exchange,提问作者user18737194
相关产品推荐
相关产品推荐

