NumPy实现自定义矩阵乘法生成指定三维数组的方法
NumPy高效实现自定义规则的三维矩阵运算
给定两个矩阵A、B,需要构造三维数组C,运算规则如下:
C[k,i,j] = A[k,j] * B[k,i]
常见的实现思路有两种:
- 手写三层Python循环逐索引遍历计算乘积:逻辑直观易懂,但纯Python循环执行效率极低,数据量稍大时运行速度很难满足需求。
- 调用
np.einsum实现:初上手时很容易写错索引对应关系,得不到正确结果,甚至会误以为该函数不适用于这类场景,实际上这是这类自定义张量运算的最优实现方式之一。
参考@Warren Weckesser 给出的方案,正确的np.einsum调用只需要一行代码,简洁且执行效率极高:
C = np.einsum('kj,ki->kij', A, B)
写法逻辑非常直接:传入的索引字符串和运算规则的下标一一对应,逗号分隔两个输入矩阵的维度顺序,箭头后标注输出数组的维度顺序即可。NumPy底层会自动完成维度匹配和广播计算,经过优化的C实现运行效率比手写Python循环高数个量级,不需要手动调整数组形状或者写复杂的广播逻辑。
内容的提问来源于stack exchange,提问作者snatchysquid
相关产品推荐
相关产品推荐

