爱因斯坦求和约定einsum:从物理符号到张量运算的编程利器
1. 爱因斯坦求和:从物理符号到编程利器的蜕变
如果你在深度学习、科学计算或者数据处理领域摸爬滚打过一阵子,大概率会见过一个看起来有点神秘的函数:einsum。无论是NumPy、PyTorch、TensorFlow还是JAX,这个函数都稳坐其中。它的全称是“爱因斯坦求和约定”(Einstein summation convention),名字听起来就充满了物理学的厚重感。但别被吓到,它本质上是一个极其强大且优雅的张量运算描述工具。简单来说,einsum允许你用一串简单的下标字符串,来定义复杂的多维数组(张量)之间的运算,比如矩阵乘法、转置、求和、对角线提取等等。它就像是一套专门为多维数据操作设计的“迷你语言”,一旦掌握,代码的简洁性和表达力会得到质的飞跃。
我最初接触einsum时,也被它那套下标规则弄得有点晕。但当我真正理解其背后的思想,并在实际项目中用它替换掉层层嵌套的循环和多个库函数调用后,我才体会到什么叫“降维打击”。它不仅让代码更清晰,避免了中间变量的创建,而且在许多框架中,底层优化做得很好,性能往往更优。这篇文章,我就以一个实践者的角度,带你彻底吃透einsum。我们会从它的思想本源讲起,拆解每一个语法细节,并通过大量实际场景的例子,让你看到它如何化繁为简。无论你是正在为复杂的张量操作头疼的研究员,还是想写出更优雅代码的工程师,相信这篇详解都能让你有所收获。
2. 核心思想:为什么爱因斯坦的“偷懒”成了程序员的福音
2.1 从求和约定到编程抽象
爱因斯坦求和约定的诞生,源于物理学家阿尔伯特·爱因斯坦的一个“偷懒”行为。在广义相对论等领域的公式中,涉及大量张量分量求和,求和符号Σ频繁出现,使得公式异常冗长。爱因斯坦发现,当某个指标在单项式的上下标中各出现一次时,就默认表示对这个指标的所有可能值求和。于是,他省去了求和符号Σ。
例如,矩阵乘法C = A · B,其分量形式为:C_{ik} = Σ_j A_{ij} * B_{jk}按照爱因斯坦约定,求和下标j在右边出现了两次(一次在A,一次在B),因此可以省略求和符号Σ,直接写作:C_{ik} = A_{ij} B_{jk}
这就是einsum的灵魂:它将这个数学约定抽象成了编程接口。我们不再需要显式地写出循环和累加,只需要告诉计算机:
- 输入张量有哪些,它们的维度用什么标签(下标)表示。
- 输出张量我们想要什么维度标签。
- 那些在输入中出现但未在输出中出现的标签,就是我们需要求和消去的维度。
这种描述方式具有声明式编程的特点:我们只关心“要做什么”(What),而不是“怎么做”(How)。具体的循环、内存排布、并行优化,都交给底层的einsum实现去处理,这通常比我们自己手写的循环高效得多。
2.2 einsum 的通用语法格式解析
几乎所有实现都遵循相似的语法。以NumPy的np.einsum为例,其基本调用形式为:
result = np.einsum(subscripts, *operands)其中,subscripts是一个字符串,定义了整个运算。operands是输入的张量。
下标字符串的规则是理解einsum的关键,我把它总结为“三步解析法”:
第一步:定义输入张量的轴标签每个输入张量用一组逗号分隔的字母序列表示。例如,‘ij, jk’表示有两个输入张量,第一个张量的两个维度分别标记为i和j,第二个张量的两个维度标记为j和k。字母通常是单个小写英文字母,但理论上可以是任何可哈希的字符。
第二步:定义输出张量的轴标签在输入描述之后,用->符号连接输出描述。例如,‘ij, jk -> ik’。这明确指出了我们想要一个维度标签为i和k的输出张量。
第三步:执行“求和约减”核心规则:所有在输入中出现,但未在输出中出现的标签,都会被求和消去。在上例‘ij, jk -> ik’中,标签j在输入中出现了两次(在第一个和第二个张量中),但在输出ik中没有出现。因此,einsum会自动对j维度进行求和。这正好完成了矩阵乘法C_{ik} = Σ_j A_{ij} * B_{jk}。
如果输出字符串被省略,比如np.einsum(‘ij, jk’, A, B),那么einsum会默认输出包含所有出现过且只出现一次的标签,并按字母表顺序排列。对于‘ij, jk’,标签i和k只出现一次,所以输出形状会是(i, k),效果等同于‘ij, jk -> ik’。但我强烈建议始终显式写出输出标签,这能让意图更清晰,避免歧义。
注意:标签重复的含义。在同一输入张量内,重复的标签表示该张量在那个维度上是对角线元素,或者要求该维度必须相等(取决于上下文)。我们会在后面的高级用法里详细讲。
3. 从基础到精通:einsum 实战示例全解
理解了语法,最好的学习方式就是看例子。我们从最简单的操作开始,逐步增加复杂度。你可以打开Python解释器,跟着一起操作。
3.1 基础单张量操作
这些操作通常可以用其他专门的函数(如sum,transpose,diag)完成,但einsum提供了一种统一的视角。
3.1.1 求和对一个二维矩阵A(形状(3, 4))进行各种求和。
import numpy as np A = np.arange(12).reshape(3, 4) # 对所有元素求和 sum_all = np.einsum(‘ij->’, A) # 等价于 A.sum() # 对第0轴(行)求和,消去i,保留j sum_axis0 = np.einsum(‘ij->j’, A) # 等价于 A.sum(axis=0) # 对第1轴(列)求和,消去j,保留i sum_axis1 = np.einsum(‘ij->i’, A) # 等价于 A.sum(axis=1)‘ij->’:输出为空,意味着消去所有维度(i和j),即全求和。‘ij->j’:输出保留j,意味着消去i,即按行求和(i是行索引),结果是一个长度为j(列数)的向量。‘ij->i’:同理,按列求和。
3.1.2 转置
A = np.arange(12).reshape(3, 4) AT = np.einsum(‘ij->ji’, A) # 等价于 A.T 或 np.transpose(A)这非常直观:我们只是交换了输出标签的顺序。
3.1.3 提取对角线对于方阵B(形状(5,5)),提取其主对角线。
B = np.arange(25).reshape(5, 5) diag = np.einsum(‘ii->i’, B) # 等价于 np.diag(B)这里,输入标签是‘ii’,这表示我们只关心i等于j的那些元素。输出为‘i’,意味着我们将这两个相同的维度压缩成一个一维数组,内容就是对角线元素。
3.1.4 逐元素乘法(哈达玛积)
A = np.arange(6).reshape(2, 3) B = np.arange(6, 12).reshape(2, 3) C = np.einsum(‘ij, ij->ij’, A, B) # 等价于 A * B输入和输出标签完全一致,表示对应位置的元素相乘,没有求和发生。
3.2 核心双张量操作
这是einsum大放异彩的地方,它能用简洁的表达式替代多个库函数调用。
3.2.1 矩阵乘法与向量内积这是einsum的“Hello World”。
# 矩阵乘法 (2,3) * (3,4) -> (2,4) A = np.random.randn(2, 3) B = np.random.randn(3, 4) C = np.einsum(‘ik, kj->ij’, A, B) # 等价于 np.dot(A, B) 或 A @ B # 注意:这里我用了i,k,j,和之前的i,j,k例子是等价的,标签名字是任意的。 # 向量内积 v1 = np.array([1, 2, 3]) v2 = np.array([4, 5, 6]) dot_product = np.einsum(‘i, i->’, v1, v2) # 等价于 np.dot(v1, v2)向量内积中,相同的标签i在输入中出现两次,输出中未出现,所以对i求和,得到一个标量。
3.2.2 张量缩并(Tensor Contraction)这是比矩阵乘法更一般的概念。例如,计算两个三维张量在特定轴上的缩并。
# 张量A形状(2,3,4), B形状(4,3,5)。对A的第2轴和B的第1轴进行缩并。 A = np.random.randn(2, 3, 4) B = np.random.randn(4, 3, 5) # 我们想消去的是A的‘k’和B的‘j’。注意我们给轴赋予了有意义的标签。 # A: 轴0->i, 轴1->j, 轴2->k # B: 轴0->k, 轴1->j, 轴2->l # 缩并发生在 (A的k, B的k) 和 (A的j, B的j)?不对,仔细看。 # 我们想用A的轴2(大小为4)与B的轴0(大小为4)做乘法求和,同时用A的轴1(大小为3)与B的轴1(大小为3)做乘法求和。 # 这意味着有两个维度要消去。正确的表达式是: C = np.einsum(‘ijk, jkl->il’, A, B) # 错误!这样只消去了j和k中的一个。 # 实际上,我们需要明确:A的第二个轴(索引1,大小3)对应B的第二个轴(索引1,大小3)。 # A的第三个轴(索引2,大小4)对应B的第一个轴(索引0,大小4)。 # 所以,设A的标签为 i,j,k; B的标签为 k,j,l。 # 那么,在运算中,j和k都出现了两次(在A和B中各一次),且输出’il’中不包含j和k。 # 因此,einsum会对j和k两个维度都进行求和。结果形状为(i, l)即(2, 5)。 C = np.einsum(‘ijk, kjl->il’, A, B) # 这才是正确的。这个例子有点绕,但它展示了einsum处理复杂缩并的能力。关键在于清晰地定义每个轴的标签,并让需要缩并的轴使用相同的标签。
3.2.3 外积
a = np.array([1, 2, 3]) b = np.array([4, 5, 6, 7]) outer = np.einsum(‘i, j->ij’, a, b) # 等价于 np.outer(a, b)输出标签‘ij’包含了输入的所有标签,且它们都只出现一次,因此没有求和。结果是一个(3, 4)的矩阵,其中outer[i, j] = a[i] * b[j]。
3.3 高级与组合操作
当操作涉及三个及以上张量时,einsum的简洁性优势更加明显。
3.3.1 批量矩阵乘法在深度学习中,我们经常处理批量数据。假设有一批矩阵A_batch(形状(batch, m, n)) 和B_batch(形状(batch, n, p)),我们需要对每一对矩阵进行乘法。
batch, m, n, p = 10, 5, 6, 7 A_batch = np.random.randn(batch, m, n) B_batch = np.random.randn(batch, n, p) # 使用einsum进行批量矩阵乘法 C_batch = np.einsum(‘bij, bjk->bik’, A_batch, B_batch) # 形状 (10, 5, 7)标签b代表批次维度,它在输入和输出中都存在,因此不参与求和,只是逐批次地进行ij, jk -> ik的矩阵乘法。这比用循环快得多,也清晰得多。
3.3.2 双线性变换这是机器学习中常见的操作,例如在注意力机制中:x^T W y,其中x和y是向量,W是矩阵。
x = np.random.randn(5) # 形状 (5,) W = np.random.randn(5, 6) # 形状 (5, 6) y = np.random.randn(6) # 形状 (6,) # 结果应为标量 result = np.einsum(‘i, ij, j->’, x, W, y) # 分解看: (i, ij) 对i求和得到中间向量 (j),再与 (j) 对j求和得到标量。 # 等价于 x.dot(W).dot(y) 或 np.dot(x, np.dot(W, y))3.3.3 张量链式乘法(多路缩并)计算像A_{ab} B_{bcd} C_{de}这样的表达式。
A = np.random.randn(4, 5) # ab B = np.random.randn(5, 3, 6) # bcd C = np.random.randn(6, 7) # de # 目标:对b和d进行求和。输出应为维度 (a, c, e) result = np.einsum(‘ab, bcd, de->ace’, A, B, C)这个表达式一次性完成了多个张量的缩并,如果用手写循环或者分步计算,会非常繁琐且容易出错。
4. 性能、优化与内存视图
4.1 einsum 的性能考量
很多人问,einsum快吗?答案是:取决于后端实现和具体操作。
在NumPy中,np.einsum本身是Python实现的,但它内部会尝试将表达式转换为高效的底层BLAS(基础线性代数子程序)调用,比如对于简单的矩阵乘法‘ij,jk->ik’,它会路由到np.dot。对于复杂的、无法映射到单一BLAS操作的表达式,它会使用自己的C语言循环内核。通常,对于简单的逐元素操作或小规模张量,einsum可能比专门的函数(如np.sum,np.transpose)稍慢,因为它有解析字符串的开销。但对于复杂的多张量缩并,它往往比手写Python循环快几个数量级,并且代码更安全。
在PyTorch和TensorFlow中,torch.einsum和tf.einsum是作为算子直接集成到计算图中的。它们会由框架的编译器(如PyTorch的TorchScript、TensorFlow的XLA)进行优化,可能融合多个操作,减少内存读写,从而获得很好的性能。在JAX中,jax.numpy.einsum可以与jax.jit无缝结合,被编译成高效的XLA代码。
实操心得:不要盲目使用
einsum。对于极其简单的操作(如单个轴求和、转置),使用专用函数(sum,T)可能更直观且微快。但对于涉及两个及以上张量、维度变换复杂的操作,einsum在可读性和性能上通常是更优选择。在性能关键路径上,建议对einsum和替代实现(如使用matmul,tensordot组合)进行简单的基准测试。
4.2 优化标志
NumPy的einsum提供了一个optimize参数,这是提升性能的大杀器。
# 未优化 result = np.einsum(‘ab, bcd, de->ace’, A, B, C) # 使用‘optimal’优化,NumPy会寻找计算路径中缩并顺序的最优解 result_opt = np.einsum(‘ab, bcd, de->ace’, A, B, C, optimize=‘optimal’)对于涉及三个及以上张量的链式乘法,计算顺序(先算哪两个)会极大影响所需的浮点运算次数(FLOPs)和中间内存占用。optimize=‘optimal’会让NumPy在计算前,使用类似动态规划算法(opt_einsum库中的算法)寻找最优或接近最优的缩并路径。对于大规模张量运算,开启优化可能带来数倍甚至数十倍的性能提升。
optimize参数也可以是‘greedy’(贪心算法,更快但可能不是最优)或True(等同于‘greedy’)。对于生产环境中的复杂einsum运算,总是设置optimize=‘optimal’是一个好习惯。
4.3 理解输出与内存视图
einsum的另一个重要特性是,它尽可能返回一个原始数组的视图而非副本,尤其是在不涉及求和的操作时。例如转置‘ij->ji’和提取对角线‘ii->i’,返回的是视图。这意味着修改返回的数组可能会影响原数组。
A = np.arange(9).reshape(3,3) A_T_view = np.einsum(‘ij->ji’, A) A_T_view[0,0] = 100 print(A[0,0]) # 输出 100, 原数组被修改了!而涉及求和的操作(如‘ij->i’),必然需要计算并分配新内存,返回的是副本。
注意事项:如果你需要一份独立的数据,记得使用
.copy()方法。例如result = np.einsum(‘ij->ji’, A).copy()。
5. 避坑指南与常见问题排查
即使理解了原理,在实际使用中还是会踩一些坑。下面是我总结的几个常见问题和解决方法。
5.1 维度不匹配错误
这是最常见的错误。einsum要求,共享的标签对应的维度大小必须相等。
A = np.ones((3, 4)) B = np.ones((5, 6)) try: C = np.einsum(‘ij, jk->ik’, A, B) # 会报错! except ValueError as e: print(e) # 很可能提示:size of dimension ‘j’ must be the same错误在于,第一个张量的j维度大小是4,而第二个张量的j维度大小是5,它们不匹配。仔细检查输入张量的形状和下标字符串的对应关系。
5.2 广播机制
NumPy的einsum支持广播,但规则需要明确。广播发生在维度标签缺失的情况下。
A = np.ones((3, 4, 5)) # 形状 (3,4,5) B = np.ones((5,)) # 形状 (5,) # 我们想用B乘以A的最后一个维度 C = np.einsum(‘ijk, k->ijk’, A, B) # 正确,B沿缺失的i和j轴广播 # 等效于 A * B (利用NumPy广播)这里,B的下标是‘k’,而A的下标是‘ijk’。在运算时,B会在i和j维度上自动广播(复制)以匹配A的形状。
再看一个更复杂的例子,涉及多个广播维度:
A = np.ones((2, 1, 3, 4)) # 形状 (2,1,3,4) B = np.ones((5, 4, 2)) # 形状 (5,4,2) # 我们希望进行某种运算,其中A的轴0对应B的轴2,A的轴3对应B的轴1。 # A的轴1(大小为1)和B的轴0(大小为5)需要广播。 # 设 A: i, j, k, l; B: m, l, i # 输出我们想要包含广播后的维度 j, k, m。 result = np.einsum(‘ijkl, mli->jkm’, A, B) # 解释: # - 标签 ‘i’ 在A和B中都出现,且未在输出出现,所以求和。对应A轴0和B轴2,大小必须相等(都是2)。 # - 标签 ‘l’ 在A和B中都出现,且未在输出出现,所以求和。对应A轴3和B轴1,大小必须相等(都是4)。 # - 标签 ‘j’ 只在A出现,在输出出现,所以保留。A轴1大小为1,在输出维度中会广播。 # - 标签 ‘k’ 只在A出现,在输出出现,所以保留。 # - 标签 ‘m’ 只在B出现,在输出出现,所以保留。B轴0大小为5。 # 最终输出形状为 (1, 3, 5) -> 广播后为 (3, 5)。因为大小为1的维度在NumPy中通常会被压缩。广播规则可以很强大,但也容易让人困惑。画一张张量形状和标签的对应图,是理清思路的好方法。
5.3 标签重复:对角线与迹运算
之前提到,在同一输入张量内重复标签,表示只取该维度索引相同的元素(对角线)。
# 创建一个三维张量 (2,3,2),其中我们希望第一个和第三个维度索引相同 # 这有点奇怪,通常用于更高维度的“对角”部分提取 A = np.random.randn(3, 4, 3) # 提取 i==k 的那些“面”,输出形状为 (3, 4) diag_slices = np.einsum(‘ijk->ij’, A) # 错误!这其实是求和了k维度。 correct_diag = np.einsum(‘iji->ij’, A) # 正确。i在第一个和第三个位置重复。更常见的例子是计算矩阵的迹(对角线元素之和):
M = np.arange(9).reshape(3,3) trace = np.einsum(‘ii->’, M) # 等价于 np.trace(M)对于‘ii->’,它先提取对角线(i==i的元素),然后因为输出为空字符串,再对这个一维对角线数组求和,得到迹。
5.4 调试技巧:使用einsum_path
对于非常复杂的表达式,如果担心性能或想了解内部的优化路径,可以使用np.einsum_path。它会返回计算路径和预估的成本。
path_info = np.einsum_path(‘ab, bcd, de->ace’, A, B, C, optimize=‘optimal’) print(path_info[0]) # 最优的缩并顺序,例如 [(0, 2), (0, 1)] print(path_info[1]) # 详细的成本信息[(0, 2), (0, 1)]表示先缩并第0个参数(A)和第2个参数(C),产生一个中间张量,然后再将这个中间张量与第1个参数(B)缩并。通过查看路径,你可以理解einsum是如何分解复杂运算的。
5.5 与tensordot,matmul的对比
numpy.tensordot(a, b, axes)是专门用于指定轴缩并的函数,功能上是einsum的子集。例如,np.tensordot(A, B, axes=([1],[0]))大致等价于np.einsum(‘ij, jk->ik’, A, B)。tensordot的接口对于简单的双张量缩并更直接,但不如einsum表达力强。
numpy.matmul和@运算符是批量矩阵乘法的标准实现,对于符合其固定模式的操作(最后两个维度做矩阵乘,前面维度广播),使用它们更符合习惯且可能经过特殊优化。einsum则提供了无与伦比的灵活性。
选择建议:
- 标准批量矩阵乘:用
@或matmul。 - 简单的双张量指定轴缩并:
tensordot也可考虑。 - 复杂的、非标准的、涉及三个及以上张量的操作:毫不犹豫地用
einsum。
6. 在深度学习框架中的应用
在PyTorch和TensorFlow中,einsum的用法与NumPy几乎一致,并且能够利用GPU加速和自动微分,在定义自定义层或损失函数时非常有用。
PyTorch 示例:实现注意力分数计算假设我们有一批查询(Query)和键(Key),计算注意力分数。
import torch batch, num_heads, seq_len, d_k = 32, 8, 10, 64 Q = torch.randn(batch, num_heads, seq_len, d_k) K = torch.randn(batch, num_heads, seq_len, d_k) # 计算 Q * K^T,对最后一个维度d_k做点积 # 期望输出形状: (batch, num_heads, seq_len, seq_len) scores = torch.einsum(‘bhid, bhjd->bhij’, Q, K) / (d_k ** 0.5) # 对比用其他方法实现,可能需要 permute 和 matmul,代码更冗长。 # scores_alt = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5) # 这个其实更直接,但einsum表达清晰。这里,einsum清晰地表达了我们想对d维度进行求和(点积),并保持其他维度不变。
TensorFlow 示例:双线性注意力
import tensorflow as tf # 假设 x: (batch, dim_x), y: (batch, dim_y), W: (dim_x, dim_y, dim_att) x = tf.random.normal((32, 50)) y = tf.random.normal((32, 60)) W = tf.random.normal((50, 60, 10)) # 计算双线性注意力分数: x^T W y -> (batch, dim_att) attention = tf.einsum(‘bi, ijo, bj->bo’, x, W, y) # 这个表达式一次性完成了复杂的双线性变换,非常简洁。JAX 示例:与 JIT 编译结合JAX的einsum可以和jax.jit完美结合,获得极致性能。
import jax.numpy as jnp from jax import jit def complex_tensor_operation(A, B, C): # 一个复杂的多张量运算 return jnp.einsum(‘ab, bcd, def, fg->ag’, A, B, C, D) # 编译这个函数 compiled_func = jit(complex_tensor_operation) # 后续调用 compiled_func 会执行编译后的高效代码7. 思维拓展:将 einsum 融入编程思维
掌握了einsum的语法后,更重要的是培养一种“einsum思维”。当你面对一个多维数组操作问题时,可以尝试以下步骤:
- 画图或标注:在白板或纸上画出每个输入张量,标出它们的维度和大小。给每个维度起一个标签(如
i, j, k, l)。 - 定义目标:明确你想要的输出张量,它的每个维度来自哪里?是保留的输入维度,还是新的维度?
- 写出下标字符串:根据输入和输出,写出
einsum表达式。思考哪些维度需要求和(消去),哪些需要保留。 - 验证形状:在写代码前,用手算或心算验证一下输出形状是否符合预期。记住规则:输出形状由输出标签决定,每个标签的大小取自任意一个包含它的输入张量的对应维度(必须一致)。
- 考虑优化:对于复杂运算,使用
optimize=‘optimal’。
我个人习惯在写涉及张量的代码时,优先考虑能否用einsum实现。它迫使你清晰地思考数据的维度流动,写出的代码往往更健壮,更易于检查。有一次,我重构了一段包含多个transpose和reshape的旧代码,用一个einsum表达式就替代了,不仅行数减少了70%,而且逻辑一目了然,还消除了一个隐蔽的维度对齐错误。
最后,再分享一个调试复杂表达式的小技巧:如果表达式出错了,可以尝试分步计算。例如,对于einsum(‘ab, bcd, de->ace’, A, B, C),可以先计算中间结果temp = einsum(‘ab, de->abde’, A, C),然后再与B运算。虽然效率不高,但能帮你理清维度是如何交互的。einsum是一个需要练习的工具,开始时可能会觉得下标游戏很烧脑,但一旦形成肌肉记忆,你就会发现它已经成为处理多维数据时不可或缺的瑞士军刀。
