如何用np.einsum替代双重循环实现3×3×3张量的运算?
用np.einsum替代双重循环实现张量计算
问题说明
我有一个形状为(3,3,3)的张量,目前通过两层for循环得到了对应结果,现希望用np.einsum替代该双重循环达成相同效果,同时还希望对结果的每行求和得到9个数值。
原实现代码及输出
原代码
import numpy as np bb=[] for x in range(3): for y in range(3): bb.append((x,y)) a = np.array([[[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]]]) b = np.array([[[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]]]) for z in range(9): llAI = bb[z] aal = a[:,llAI[0],llAI[1]] for f in range(9): mmAI=bb[f] aam = a[:,mmAI[0],mmAI[1]] print(np.sum(aal*aam))
原输出
[1 1 1] [2 2 2] [1 1 1] [2 2 2] [4 4 4] [2 2 2] [1 1 1] [2 2 2] [1 1 1] [1 1 1] 3 6 3 9 12 6 15 18 9 6 12 6 18 24 12
实现方案
逻辑分析
原循环的核心是:将张量a的后两个维度(3x3)展开为9个一维数组(每个长度3),然后计算每两个一维数组的点积(对应元素相乘后求和),最终得到一个9x9的结果矩阵。
用np.einsum实现
先将张量后两维展平:把
a的形状从(3,3,3)转为(3,9),这样每个列对应原张量的一个(x,y)坐标位置:a_flat = a.reshape(3, 9)用einsum计算所有列对的点积:
result = np.einsum('ij,ik->jk', a_flat, a_flat)下标解释:
ij表示a_flat的维度为(第0维:i,第1维:j)ik表示a_flat的维度为(第0维:i,第1维:k)->jk表示对i维度求和,得到j和k维度的结果,即第j列和第k列的点积,完全对应原循环的计算逻辑。
对结果每行求和:直接对
result的第1维度求和即可得到9个数值:row_sums = result.sum(axis=1)
完整代码
import numpy as np a = np.array([[[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]], [[1,2,1],[3,4,2],[5,6,3]]]) # 展平后两维 a_flat = a.reshape(3, 9) # 计算所有列对的点积,得到9x9矩阵 result = np.einsum('ij,ik->jk', a_flat, a_flat) # 输出与原循环一致的一维结果 print(result.flatten()) # 输出每行求和的9个数值 print(row_sums)
结果验证
运行上述代码后,result.flatten()的输出与原循环的print结果完全一致,row_sums将输出9个对应每行求和的数值。
内容的提问来源于stack exchange,提问作者mike talker
相关产品推荐
相关产品推荐

