拓冰建站拓冰建站
首页 / 资讯中心 / 正文

爱因斯坦求和约定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(axis0) # 对第1轴列求和消去j保留i sum_axis1 np.einsum(‘ij-i’, A) # 等价于 A.sum(axis1)‘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的轴2A的轴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) # 提取 ik 的那些“面”输出形状为 (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-’它先提取对角线ii的元素然后因为输出为空字符串再对这个一维对角线数组求和得到迹。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是一个需要练习的工具开始时可能会觉得下标游戏很烧脑但一旦形成肌肉记忆你就会发现它已经成为处理多维数据时不可或缺的瑞士军刀。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门