ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

从爱因斯坦求和到多维数组运算:einsum在NumPy、PyTorch与TensorFlow中的核心应用

2026/8/4 3:56:37 拓冰建站 浏览量
从爱因斯坦求和到多维数组运算:einsum在NumPy、PyTorch与TensorFlow中的核心应用 1. 从“天书”到利器为什么einsum值得你花时间第一次看到einsum爱因斯坦求和约定的表达式比如np.einsum(‘ij,jk-ik’, A, B)很多人的反应和我当初一样这写的什么玩意儿一堆字母下标箭头飞来飞去看起来像某种神秘的数学咒语。但当你真正搞懂它之后会发现它可能是处理多维数组运算时最清晰、最强大、也最优雅的工具之一没有“之一”。简单来说einsum是一种通过下标标记法来定义张量多维数组运算的规则。它得名于物理学家阿尔伯特·爱因斯坦他在广义相对论中引入这套约定来简化冗长的求和公式。在编程世界尤其是NumPy、PyTorch、TensorFlow这些以数组计算为核心的库中einsum将这套思想发扬光大让你可以用一个简洁的字符串替代一系列复杂的transpose转置、reshape重塑、sum求和和dot点积操作。它解决了什么痛点想象一下你需要计算两个三维张量在特定维度上的乘积并求和或者进行复杂的张量缩并。用传统的数组方法你可能需要写好几行代码反复调整轴顺序还得小心翼翼确保维度对齐。而einsum一行就能说清楚“我要对这些下标进行求和并得到那样的输出形状。” 意图直接代码自文档化极大减少了思维负担和出错概率。无论你是做数据科学、机器学习、物理模拟还是深度学习只要涉及多维数组操作einsum都是一个绕不开的高效工具。2. einsum的核心语法读懂下标语言einsum的魔力全部浓缩在那个小小的下标字符串里。它的通用形式是np.einsum(subscripts, *operands)。其中subscripts是核心定义了整个运算的蓝图。2.1 下标字符串的构成规则下标字符串由三部分组成用逗号分隔输入操作数用箭头-指向输出。基本格式为[输入1下标],[输入2下标],...-[输出下标]。下标字符的约定每个下标是一个或多个小写字母如i,j,k每个字母代表一个维度。重复的字母意味着求和缩并。这是爱因斯坦求和约定的精髓。箭头-右侧的输出下标定义了最终结果的维度顺序和保留哪些轴。让我们从一个最简单的例子看起向量点积。向量a和b都是一维数组。import numpy as np a np.array([1, 2, 3]) b np.array([4, 5, 6])传统做法是np.dot(a, b)。用einsum怎么写np.einsum(‘i,i-’, a, b)。i,i 第一个i对应a的维度第二个i对应b的维度。两个下标都是i意味着我们要对i这个维度进行逐元素相乘并求和。- 箭头后面是空的。这表示求和后这个i维度被“消灭”了结果是一个标量0维数组。计算过程等价于sum 1*4 2*5 3*6 32。再看矩阵乘法这是einsum最经典的用例。矩阵A(2x3) 和矩阵B(3x4) 相乘。A np.random.rand(2, 3) B np.random.rand(3, 4)传统做法是np.matmul(A, B)。einsum表达式为np.einsum(‘ij,jk-ik’, A, B)。ij 对应Ai是第0轴行j是第1轴列。jk 对应Bj是第0轴行k是第1轴列。-ik 输出下标是i和k。注意下标j出现在了输入中但没有出现在输出中。根据规则所有在输入中出现但未在输出中出现的下标都会被求和缩并。所以这里是对j维度进行求和。计算过程对于输出结果的每一个位置(i, k)其值等于sum_over_j( A[i, j] * B[j, k] )。这正是矩阵乘法的定义。2.2 输出下标的控制艺术输出下标是你控制结果形态的遥控器。通过精心设计输出下标你可以实现转置、取对角线、广播等操作。1. 显式指定输出顺序实现转置假设有一个矩阵M(3x4)我们想将其转置。传统做法是M.T或np.transpose(M)。 用einsumnp.einsum(‘ij-ji’, M)。输入下标ij输出下标ji。这直接告诉程序“把i轴和j轴交换位置。” 意图一目了然。2. 保留求和轴实现按行/列求和并保持维度对矩阵M按列求和即对行轴i求和通常得到一行向量。传统做法M.sum(axis0)得到形状(4,)。 如果想保持二维结构得到一个(1, 4)的形状可能需要M.sum(axis0, keepdimsTrue)。 用einsum可以更直观np.einsum(‘ij-j’, M)等价于sum(axis0)。 如果想显式保留那个被求和的维度尽管它长度为1在输出下标中省略即可但einsum本身不直接支持keepdims。不过你可以通过添加一个虚拟维度来实现类似效果但这通常不如keepdims直接。这里更展示einsum在求和维度的控制上是“全有或全无”的。3. 取矩阵的迹对角线元素之和对于方阵S(n x n)迹是np.trace(S)。 用einsumnp.einsum(‘ii-’, S)。输入下标ii表示两个维度是同一个i。这代表我们只取ii的元素即对角线。输出为空表示对这些对角线元素求和得到一个标量。4. 提取对角线元素如果不想求和只想取出对角线元素组成一个向量呢 用einsumnp.einsum(‘ii-i’, S)。输入ii仍然表示取对角线。输出i表示将结果沿着i这个维度排列形成一个一维向量。注意下标字母的语义是局部的。np.einsum(‘ij,jk-ik’, A, B)和np.einsum(‘ab,bc-ac’, A, B)是完全等价的。字母本身没有特定含义它只是在你定义的这次运算中用来标记维度的临时标签。这给了你很大的灵活性但也要求你在一个表达式内部保持一致性。3. 进阶应用解锁多维张量操作einsum的真正威力在处理三维及以上的张量时才会完全展现。很多用传统方法写起来很拧巴的操作用einsum可以优雅地一行搞定。3.1 张量缩并高维空间的“点积”张量缩并是向量点积和矩阵乘法在高维空间的推广。例如有一个三维张量T(2x3x4) 和一个二维矩阵M(4x5)我们想在第3轴T的最后一个轴和第0轴M的第一个轴上进行缩并。T np.random.rand(2, 3, 4) M np.random.rand(4, 5)我们想要的结果形状是 (2, 3, 5)。传统方法可能需要np.tensordot(T, M, axes([2], [0]))。 用einsum一目了然result np.einsum(‘ijk,kl-ijl’, T, M)。ijk对应T的三个轴。kl对应M的两个轴。下标k同时出现在T和M中但未出现在输出ijl中因此对k轴进行求和缩并。输出ijl决定了结果的形状和轴顺序i(2),j(3),l(5)。3.2 批量矩阵乘法在深度学习中我们经常遇到批量数据。例如有一批矩阵A_batch(batch_size, m, n) 和B_batch(batch_size, n, p)我们需要对每一对矩阵进行独立的乘法。 传统方法可能要用循环或者使用np.matmul它天然支持批量操作。 用einsum可以清晰地表达np.einsum(‘bij,bjk-bik’, A_batch, B_batch)。下标b代表了批量维度。这个表达式明确告诉我们“在保持批量b独立的前提下对每一对(i,j)和(j,k)矩阵执行ij,jk-ik的乘法。”这比思考np.matmul的轴对齐规则要直观得多尤其是当维度更多更复杂时。3.3 外积与广播向量的外积np.outer(a, b)生成一个矩阵其中M[i,j] a[i] * b[j]。 用einsumnp.einsum(‘i,j-ij’, a, b)。输入下标i和j没有重复意味着不发生求和。输出下标ij意味着将i和j组合成一个二维输出。这本质上是将两个一维数组通过广播机制进行相乘。广播机制在einsum中隐式工作。只要维度能通过广播对齐且符合下标规则就可以运算。例如将一个向量v(n,) 加到矩阵M(m, n) 的每一行 传统做法M v(利用NumPy广播)。einsum做法np.einsum(‘ij,j-ij’, M, v)不对这样会触发对j的求和。我们不想求和只想广播。 其实对于这种纯粹的、带广播的元素级运算einsum并非最佳选择直接用M v更清晰。einsum更擅长涉及求和缩并的运算。对于广播加法用einsum需要一点技巧np.einsum(‘ij,j-ij’, M, v) - np.einsum(‘ij,j-ij’, M, v)? 这显然不对。正确的理解是einsum的广播发生在“未标记的维度”上但它的规则更侧重于下标指定的维度关系。对于简单的广播加法直接使用数组运算即可不必强行使用einsum。3.4 复杂示例注意力机制中的计算Transformer模型中的缩放点积注意力其核心计算是Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V。 假设Q,K,V都是三维张量形状为 (batch_size, seq_len, d_model)。 用einsum可以非常优雅地写出import numpy as np # 假设维度 batch (b), 序列长度 (s), 键/值维度 (d_k), 查询维度 (d_q), 值维度 (d_v) # 通常 d_q d_k b, s, d_k, d_v 10, 20, 64, 64 Q np.random.randn(b, s, d_k) K np.random.randn(b, s, d_k) V np.random.randn(b, s, d_v) # 计算 QK^T 对每个batch和每个查询-键对进行点积 # 输出形状应为 (b, s, s) attn_scores np.einsum(‘bqd,bkd-bqk’, Q, K) # 这里 q 和 k 都代表 seq_len但用不同字母区分位置 # 等价于对于每个batch b每个查询位置 q每个键位置 k计算 Q[b,q,:] 和 K[b,k,:] 的点积。 # 缩放和softmax scaled_scores attn_scores / np.sqrt(d_k) attn_weights np.exp(scaled_scores) / np.exp(scaled_scores).sum(axis-1, keepdimsTrue) # 计算加权和 # attn_weights 形状 (b, s, s), V 形状 (b, s, d_v) # 对键位置 k 进行求和得到每个查询位置 q 对应的上下文向量 context np.einsum(‘bqk,bkv-bqv’, attn_weights, V)这段代码清晰地展示了信息流动Q和K相互作用产生注意力权重权重再作用于V。下标bqd,bkd-bqk和bqk,bkv-bqv完美刻画了张量间的缩并关系比用一堆transpose和matmul要直观太多。4. 性能、优化与内存细节虽然einsum在表达上非常简洁但它的性能并非总是最优。理解其背后的机制有助于你在正确的地方使用它。4.1 einsum的执行路径与优化当你调用np.einsum时NumPy 内部会经历以下几个步骤解析下标字符串解析输入输出下标确定哪些维度需要求和以及结果的形状。路径优化对于涉及多个操作数两个以上的复杂einsum求和顺序对性能影响巨大。例如计算np.einsum(‘ij,jk,kl-il’, A, B, C)可以先算(A*B)再乘C也可以先算(B*C)再与A乘。不同的顺序产生的中间数组大小不同计算量和内存占用差异显著。NumPy 的einsum从某个版本开始会尝试使用一种类似“动态规划”的算法来寻找最优或近似最优的求和路径收缩顺序。你可以通过np.einsum_path函数来查看它选择的路径和预估的成本。path_info np.einsum_path(‘ij,jk,kl-il’, A, B, C, optimize‘optimal’) print(path_info[0]) # 显示计算路径 print(path_info[1]) # 显示详细成本信息optimize参数是关键。optimizeFalse禁用优化按直观顺序计算可能很慢。optimizeTrue默认会尝试寻找较优路径。optimize‘optimal’会寻找理论最优路径但对于大量操作数可能搜索较慢。optimize‘greedy’是默认的启发式算法在速度和效果间取得平衡。执行计算根据优化后的路径调用底层通常是BLAS库的矩阵运算或执行循环进行求和。4.2 与专用函数的性能对比对于常见的、有专用函数的操作直接调用专用函数通常更快因为它们是高度优化的。矩阵乘法np.dot(A, B)或np.matmul(A, B)通常比np.einsum(‘ij,jk-ik’, A, B)稍快因为它们直接调用高度优化的BLAS库如OpenBLAS, MKL。einsum需要经过一层解析和分发。转置A.T或A.transpose()是视图操作几乎零成本。np.einsum(‘ij-ji’, A)会创建一个新的数组有内存分配和复制开销。迹np.trace(A)和np.einsum(‘ii-’, A)性能接近但专用函数可能略有优势。点积np.dot(a, b)和np.einsum(‘i,i-’, a, b)性能类似。那么什么时候该用einsum操作复杂没有现成的专用函数这是einsum的主场。比如前面提到的张量缩并、批量特定维度乘法等。代码可读性优先即使有替代方案如果einsum表达式能一眼看清计算意图如注意力机制为了代码的清晰和可维护性牺牲一点微不足道的性能是值得的。原型设计和探索在尝试新的数学公式或模型结构时用einsum快速验证想法非常方便。实操心得路径优化是双刃剑。对于非常复杂的表达式如涉及4个以上张量开启optimizeTrue能带来数量级的性能提升。但优化过程本身有开销。对于在循环中反复调用的、非常简单的einsum如简单的矩阵乘法关闭优化optimizeFalse有时反而更快因为避免了每次调用时的路径分析开销。我的经验法则是在循环外预先计算einsum_path然后在循环内使用固定的路径或者对于简单操作直接使用专用函数。4.3 内存占用考量einsum在计算过程中可能会产生巨大的中间数组。例如计算三个大矩阵的乘积einsum(‘ab,bc,cd-ad’, A, B, C)。如果按照(A*B)*C的顺序会先产生一个形状为(a, c)的中间数组再与C乘。如果a, b, c, d都很大这个中间数组可能耗尽内存。np.einsum_path提供的优化一个重要目标就是最小化中间数组的大小。它会评估不同收缩顺序下中间结果的最大体积。在内存紧张的情况下务必使用einsum_path检查并选择内存友好的路径。有时手动将一个大einsum拆分成多个步骤并适时使用del释放中间变量是更稳妥的做法。5. 跨框架的einsumNumPy, PyTorch, TensorFloweinsum的概念已被主流深度学习框架广泛采纳语法几乎完全一致这带来了极大的便利。NumPy:np.einsum(subscripts, *operands, outNone, dtypeNone, order‘K’, casting‘safe’, optimizeFalse)PyTorch:torch.einsum(equation, *operands)。PyTorch 的einsum支持自动微分可以无缝嵌入神经网络中。在GPU上它能调用优化的CUDA内核。TensorFlow:tf.einsum(equation, *inputs, **kwargs)。同样支持GPU和自动微分。框架间的重要差异优化策略NumPy的optimize参数在PyTorch和TensorFlow中不一定有完全相同的实现或默认行为。PyTorch的einsum底层会尝试将操作映射到一系列基础的mm,bmm,sum等操作上。TensorFlow 的tf.einsum会尝试使用MatMul等核心操作。广播规则虽然都支持广播但细微规则可能略有不同在编写跨框架代码时需要注意。性能对于能在底层映射到高度优化算子如torch.bmm,tf.matmul的einsum表达式框架的einsum性能可能接近专用函数。但对于非常特殊的缩并可能退化为通用的、较慢的实现。动态形状在PyTorch和TensorFlow的图模式下如tf.function, TorchScripteinsum对动态形状的支持可能不如NumPy灵活。编写可移植的einsum代码建议尽量使用最简单、最标准的表达式。复杂的、依赖特定优化路径的表达式在不同框架间可能性能差异大。对于性能关键的、且在各框架中都有专用函数的操作如批量矩阵乘bmm在最终部署的代码中可以考虑替换为专用函数调用以获取最佳性能。用einsum做原型验证和文档说明。测试时除了验证结果正确也关注在不同框架下的内存和速度表现。6. 常见陷阱、调试技巧与最佳实践即使理解了语法实际使用中还是会踩坑。下面是一些常见问题和解决方法。6.1 下标错误与维度不匹配这是最常遇到的问题。错误信息通常很直接但需要会解读。ValueError: operands could not be broadcast together with remapped shapes这通常意味着下标字符串暗示的维度不匹配。例如np.einsum(‘ij,jk-ik’, A, B)但A.shape[1] ! B.shape[0]。仔细检查每个操作数对应下标的维度长度是否一致或满足广播条件。ValueError: einstein sum subscripts string contains too many subscripts for operand操作数的维度数量少于下标字符串分配给它的字母数量。例如A是二维矩阵你却写了np.einsum(‘ijk’, A)。ValueError: output has more dimensions than subscripts given in einstein sum, but no ‘…’你使用了省略号...但可能用法不对或者输出下标指定的维度数与结果的实际维度数不符。调试技巧画图在纸上画出每个张量的方块图用箭头连接需要求和缩并的维度。这能直观地检查维度是否对齐。分步验证对于复杂的表达式拆分成多个简单的einsum或使用einsum_path查看中间步骤的形状。使用np.einsum_path即使不关心性能用einsum_path也能帮你确认NumPy是如何理解你的表达式的它会打印出每个收缩步骤和中间结果的形状。6.2 省略号...的使用...Ellipsis用于表示“所有其他未指定的维度”。这在处理批量数据或高维张量时非常有用可以避免写出很长一串下标。 例如有一个四维张量T(batch, channel, height, width) (b, c, h, w)我们想对每个样本、每个通道的空间位置h, w求和得到 (b, c) 的输出。 传统方法T.sum(axis(2,3))。 用einsum且不用省略号np.einsum(‘bchw-bc’, T)。 用einsum使用省略号np.einsum(‘...hw-...’, T)。 后者的好处是即使T的维度前面增加了比如多了个时间步维度t变成 (t, b, c, h, w)表达式‘...hw-...’依然适用它会自动将t, b, c视为“其他维度”并保留。这增加了代码的鲁棒性。使用省略号的规则每个操作数中最多只能有一个...。输出中可以包含...表示保留输入中对应...所代表的所有维度。...所代表的维度集合在所有输入操作数中必须能够广播对齐。6.3 数据类型与溢出einsum默认会遵循NumPy的类型提升规则。如果操作数是整数类型求和可能导致溢出。import numpy as np a np.array([100, 200], dtypenp.int8) b np.array([100, 200], dtypenp.int8) # 点积结果应为 100*100 200*200 10000 40000 50000 result np.einsum(‘i,i-’, a, b) # 可能发生溢出得到错误结果对于可能的大数求和建议先将数组转换为浮点型或更高精度的整数类型np.einsum(‘i,i-’, a.astype(np.int64), b.astype(np.int64))。6.4 最佳实践总结从简单开始先用einsum实现你熟悉的操作如点积、矩阵乘确保理解正确。下标命名有意义虽然字母是任意的但使用有助记忆的字母如batch,channel,height,width,input,output能极大提升代码可读性。优先使用专用函数对于dot,matmul,trace,transpose等简单操作直接调用专用函数通常更优。复杂表达式先优化对于涉及三个及以上操作数的运算务必使用np.einsum_path检查优化路径特别是当数据量较大时。关注内存留意路径优化报告中的“最大中间大小”确保其不会超出可用内存。测试与验证用随机数据和小规模数据验证einsum表达式的结果是否正确可以对比使用循环实现的“朴素”版本的结果。文档化在复杂的einsum表达式旁添加注释说明每个下标字母的含义和运算的物理意义。einsum是一个需要稍加练习才能熟练掌握的工具但一旦掌握它就会成为你处理多维数组运算的思维语言。它强迫你清晰地思考每个维度的去向这种清晰性本身就能减少bug。下次当你面对一堆transpose和reshape感到头晕时不妨试试用einsum来重新表述你的问题很可能你会发现一条更清晰的道路。