如何用np.einsum单步运算验证矩阵乘法结合律?
用np.einsum实现矩阵乘法结合律验证(单步运算)
1. 用单个np.einsum完成D和E的计算
你原来的D = (A@B)@C和E = A@(B@C),完全可以用单个np.einsum直接计算,无需中间调用dot:
对应D的单步einsum
D = np.einsum('il, lj, jk -> ik', A, B, C)
索引逻辑:
A的维度为i(行)×l(列)B的维度为l(行)×j(列)C的维度为j(行)×k(列)- 收缩中间的
l和j维度,最终得到i×k的结果,和(A@B)@C完全等价。
对应E的单步einsum
E = np.einsum('il, lj, jk -> ik', A, B, C)
表达式和D完全一致——这正是矩阵乘法结合律的体现:(A@B)@C = A@(B@C)。einsum直接计算三者的张量收缩时,不需要区分结合顺序,结果自然一致。
如果想更直观对应E的计算逻辑(先算B@C再乘A),也可以拆分为两步einsum,但这不属于单步运算:
E = np.einsum('il, lk -> ik', A, np.einsum('lj, jk -> lk', B, C))
完整验证代码:
import numpy as np A = np.array([[1, 1, 1], [2, 2, 2], [5, 5, 5]]) B = np.array([[0, 1, 0], [1, 1, 0], [1, 1, 1]]) C = np.array([[ 6, 4, 2], [-2, 0, 2], [ 3, 2, 1]]) # 单步einsum计算D D = np.einsum('il, lj, jk -> ik', A, B, C) # 单步einsum计算E E = np.einsum('il, lj, jk -> ik', A, B, C) print((D == E).all()) # 输出True
2. 这种方式是否为最优方案?
是的,单步einsum运算有两个核心优势:
- 内存效率更高:避免存储中间矩阵(比如
A@B或B@C),大矩阵场景下能显著减少内存占用。 - 计算效率可控:numpy的
einsum支持通过optimize=True自动选择最优收缩路径,性能可接近专门优化的矩阵乘法实现,比分步计算省去了中间数组的创建和销毁开销。
3. 多矩阵相乘的einsum应用思路
用einsum实现多矩阵相乘的核心是索引匹配:
- 给每个矩阵的维度分配唯一索引符号,相邻矩阵的公共维度(需收缩的维度)用相同符号。
- 最终结果保留你需要的索引符号,其余公共维度会自动收缩。
比如三个矩阵相乘A@B@C,只要保证A的列索引和B的行索引一致,B的列索引和C的行索引一致,最终保留A的行索引和C的列索引即可,也就是'il, lj, jk -> ik'。
内容的提问来源于stack exchange,提问作者Martin H
相关产品推荐
相关产品推荐

