从‘看不懂’到‘真香’:用乐高积木思维彻底掌握NumPy einsum

想象你面前摆着一盒五颜六色的乐高积木,每个积木块上贴着不同的标签。你的任务是根据一张神秘的说明书,把这些散落的零件组装成指定的模型——这就是einsum函数在NumPy中扮演的角色。与传统的数学符号不同,我们将用这套乐高比喻贯穿全文,让你在30分钟内建立起对einsum的直觉理解。

1. 乐高说明书:einsum符号系统解析

当我们拿到一盒乐高时,说明书会明确标注需要哪些零件(ij,jk)以及最终成品的样子(->ik)。einsum的表达式就像这样一张组装说明书:

C = np.einsum('ij,jk->ik', A, B)
  • 零件清单(->左侧):每个逗号分隔的部分代表一个输入数组的维度标签,就像列出需要的乐高零件类型和数量。ij表示需要A数组的i行j列零件,jk表示需要B数组的j行k列零件。
  • 组装规则(重复字母):当同一个标签(如j)出现在多个零件中时,表示这些零件需要按该维度拼接,就像乐高积木的凸起和凹槽必须匹配才能连接。
  • 成品规格(->右侧):描述最终输出的形状。如果某个零件标签(如j)没有出现在成品中,说明这些零件需要在该维度上"粘合固定"(求和操作)。

让我们用实际积木演示这个比喻。假设:

A = np.array([[红, 红, 红],    # 3块红色积木组成的行
              [蓝, 蓝, 蓝]])   # 3块蓝色积木组成的行 
B = np.array([[凸, 平],        # 2种连接件
              [凹, 凸],
              [平, 凹]])

执行np.einsum('ij,jk->ik', A, B)就像:

  1. 选取A的第i行j列积木和B的第j行k列连接件
  2. 检查连接件类型是否匹配(j维度相同)
  3. 将匹配的积木与连接件组合后,统计每种颜色(i)与最终连接类型(k)的组合数量

2. 组装技巧:常见操作的可视化拆解

2.1 矩阵乘法:积木墙的搭建

传统矩阵乘法就像用固定模版组装乐高墙,而einsum则像自由组合:

# 传统方法
C = np.dot(A, B)

# einsum方法
C = np.einsum('ij,jk->ik', A, B) 

用乐高术语解释:

  • ij,jk->ik:将i颜色的积木与j型连接件组合,再与k型底座连接,最终得到i颜色k底座的成品
  • ij,jk->ijk:如果保留所有标签,会得到所有中间组合(不求和)
  • ij,jk->:如果省略所有标签,则计算所有积木连接后的总量

2.2 求和操作:积木的合并统计

当某个维度标签从输出中消失时,就像把特定类型的积木打包统计:

操作示例 乐高比喻 数学含义
'ij->i' 统计每种颜色(i)的积木总数 按行求和
'ij->j' 统计每种连接件(j)的使用总量 按列求和
'ij->' 计算所有积木的总数 全矩阵求和
# 计算矩阵对角线(提取特定连接方式)
diag = np.einsum('ii->i', A)  # 只保留行列编号相同的积木

# 计算矩阵迹(统计特定连接总量)
trace = np.einsum('ii->', A)  # 计算对角线积木的总和

2.3 转置与轴交换:改变说明书顺序

调整标签顺序就像修改组装步骤:

# 常规转置
B_T = np.einsum('ij->ji', B) 

# 特定轴交换
C = np.einsum('ijk->kij', D)  # 将第3个轴移到最前

这相当于:

  1. 原说明书:先选颜色(i),再选连接件(j)
  2. 新说明书:先选连接件(j),再选颜色(i)

3. 高级搭建:多维张量操作

当处理3D或更高维张量时,einsum就像组装复杂乐高模型:

# 批量矩阵乘法 (3D张量)
result = np.einsum('bij,bjk->bik', batch_A, batch_B)
  • b:代表不同的乐高模型编号
  • ij/jk:每个模型内部的零件组装规则
  • ->bik:输出保持批次维度,内部是ik结构

对于更高维数据,可以使用省略号:

# 处理任意维度数组的后两维
output = np.einsum('...ij,...jk->...ik', tensor_A, tensor_B)

这表示:

  • ...:代表任意数量的前置维度(不参与当前操作)
  • ij,jk->ik:只对最后两个维度进行矩阵乘法

4. 避坑指南:乐高大师的实践经验

在实际使用einsum时,有几个常见问题需要注意:

数据类型陷阱:einsum不会自动提升数据类型,使用小类型可能导致溢出

a = np.ones(300, dtype=np.int8)
np.einsum('i->', a)  # 可能得到错误结果44

性能考虑:

  • 对于简单操作(如矩阵乘法),专用函数(如dot)可能更快
  • einsum的优势在于复杂操作的表达清晰度

可读性技巧:

  1. 为维度标签赋予语义(如't'代表时间,'c'代表通道)
  2. 复杂操作分步编写和验证
  3. 使用注释说明每个标签的含义
# 清晰的标签命名示例
# b: batch, c: channel, h: height, w: width
output = np.einsum('bchw,bcHW->bhwHW', input, kernel)

掌握了这套乐高思维后,你会发现自己能"看图说话"——看到数学公式就能写出对应的einsum表达式。试着用这个比喻重新理解之前的困惑点,你会发现那些神秘的符号突然变得直观起来。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐