numpy tensordot运算结果shape推导及运行原理解惑求助
Numpy tensordot 运算规则通俗讲解
核心通用规则
tensordot本质是矩阵点积向高维张量的推广,运算逻辑和结果shape推导可以按三步走:
- 第一步匹配收缩轴:参数
axes=([A要收缩的轴索引列表], [B要收缩的轴索引列表]),要求两个列表里对应位置的轴长度完全相等,否则会报错,这些轴会在运算后消失 - 第二步提取剩余轴:分别把A、B的shape里去掉要收缩的轴,剩余维度按原来的顺序保留
- 第三步拼接得到结果shape:把A剩余的维度放在前面,B剩余的维度放在后面,直接拼接就是最终输出的shape
运算的数值计算逻辑和矩阵乘法完全一致:收缩轴上对应位置的元素先相乘,再全部求和,得到结果张量对应位置的值。
第一个示例推导
已知条件
a = np.arange(60.).reshape(3,4,5) # a.shape = (3,4,5) b = np.arange(24.).reshape(4,3,2) # b.shape = (4,3,2) c = np.tensordot(a,b, axes=([1,0],[0,1]))
推导过程
- 确定收缩轴:a要收缩的轴是索引
1(长度4)、0(长度3);b要收缩的轴是索引0(长度4)、1(长度3),长度完全匹配 - 提取剩余轴:a去掉轴1、0后,剩余的轴是索引
2,长度为5;b去掉轴0、1后,剩余的轴是索引2,长度为2 - 拼接得到shape:
(5,)+(2,)=(5,2),和运行结果一致
运算过程说明
结果中任意位置c[i][j]的数值,等于a[:,:,i]和b[:,:,j]两个矩阵的所有元素逐一相乘后的总和。
第二个示例推导
已知条件
A = np.random.randint(2, size=(2, 6, 5)) # A.shape = (2,6,5) B = np.random.randint(2, size=(3, 2, 4)) # B.shape = (3,2,4) np.tensordot(A, B, axes=((0),(1))).shape
推导过程
- 确定收缩轴:A要收缩的轴是索引
0(长度2);B要收缩的轴是索引1(长度2),长度匹配 - 提取剩余轴:A去掉轴0后,剩余的轴是索引
1(长度6)、2(长度5),顺序为(6,5);B去掉轴1后,剩余的轴是索引0(长度3)、2(长度4),顺序为(3,4) - 拼接得到shape:
(6,5)+(3,4)=(6,5,3,4),和运行结果一致
运算过程说明
结果中任意位置c[i][j][k][l]的数值,等于A[:,i,j]和B[k,:,l]两个长度为2的向量的点积(对应位置相乘再求和)。
内容的提问来源于stack exchange,提问作者Takuya2412
相关产品推荐
相关产品推荐

