当前位置: 首页 > news >正文

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

1. 从“天书”到利器:为什么einsum值得你花时间

第一次看到einsum(爱因斯坦求和约定)的表达式,比如np.einsum(‘ij,jk->ik’, A, B),很多人的反应和我当初一样:这写的什么玩意儿?一堆字母下标,箭头飞来飞去,看起来像某种神秘的数学咒语。但当你真正搞懂它之后,会发现它可能是处理多维数组运算时,最清晰、最强大、也最优雅的工具之一,没有“之一”。

简单来说,einsum是一种通过下标标记法来定义张量(多维数组)运算的规则。它得名于物理学家阿尔伯特·爱因斯坦,他在广义相对论中引入这套约定来简化冗长的求和公式。在编程世界,尤其是NumPyPyTorchTensorFlow这些以数组计算为核心的库中,einsum将这套思想发扬光大,让你可以用一个简洁的字符串,替代一系列复杂的transpose(转置)、reshape(重塑)、sum(求和)和dot(点积)操作。

它解决了什么痛点?想象一下,你需要计算两个三维张量在特定维度上的乘积并求和,或者进行复杂的张量缩并。用传统的数组方法,你可能需要写好几行代码,反复调整轴顺序,还得小心翼翼确保维度对齐。而einsum一行就能说清楚:“我要对这些下标进行求和,并得到那样的输出形状。” 意图直接,代码自文档化,极大减少了思维负担和出错概率。无论你是做数据科学、机器学习、物理模拟还是深度学习,只要涉及多维数组操作,einsum都是一个绕不开的高效工具。

2. einsum的核心语法:读懂下标语言

einsum的魔力全部浓缩在那个小小的下标字符串里。它的通用形式是:np.einsum(subscripts, *operands)。其中subscripts是核心,定义了整个运算的蓝图。

2.1 下标字符串的构成规则

下标字符串由三部分组成,用逗号分隔输入操作数,用箭头->指向输出。基本格式为:[输入1下标],[输入2下标],...->[输出下标]

下标字符的约定:

  • 每个下标是一个或多个小写字母(如i,j,k),每个字母代表一个维度。
  • 重复的字母意味着求和(缩并)。这是爱因斯坦求和约定的精髓。
  • 箭头->右侧的输出下标,定义了最终结果的维度顺序和保留哪些轴。

让我们从一个最简单的例子看起:向量点积。向量ab都是一维数组。

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: 输出下标是ik。注意,下标j出现在了输入中,但没有出现在输出中。根据规则,所有在输入中出现但未在输出中出现的下标,都会被求和(缩并)。所以,这里是对j维度进行求和。
  • 计算过程:对于输出结果的每一个位置(i, k),其值等于sum_over_j( A[i, j] * B[j, k] )。这正是矩阵乘法的定义。

2.2 输出下标的控制艺术

输出下标是你控制结果形态的遥控器。通过精心设计输出下标,你可以实现转置、取对角线、广播等操作。

1. 显式指定输出顺序(实现转置):假设有一个矩阵M(3x4),我们想将其转置。传统做法是M.Tnp.transpose(M)。 用einsumnp.einsum(‘ij->ji’, M)

  • 输入下标ij,输出下标ji。这直接告诉程序:“把i轴和j轴交换位置。” 意图一目了然。

2. 保留求和轴(实现按行/列求和并保持维度):对矩阵M按列求和(即对行轴i求和),通常得到一行向量。传统做法M.sum(axis=0)得到形状(4,)。 如果想保持二维结构,得到一个(1, 4)的形状,可能需要M.sum(axis=0, keepdims=True)。 用einsum可以更直观:np.einsum(‘ij->j’, M)等价于sum(axis=0)。 如果想显式保留那个被求和的维度(尽管它长度为1),在输出下标中省略即可,但einsum本身不直接支持keepdims。不过,你可以通过添加一个虚拟维度来实现类似效果,但这通常不如keepdims直接。这里更展示einsum在求和维度的控制上是“全有或全无”的。

3. 取矩阵的迹(对角线元素之和):对于方阵S(n x n),迹是np.trace(S)。 用einsumnp.einsum(‘ii->’, S)

  • 输入下标ii,表示两个维度是同一个i。这代表我们只取i==i的元素,即对角线。
  • 输出为空,表示对这些对角线元素求和,得到一个标量。

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同时出现在TM中,但未出现在输出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)

  • 输入下标ij没有重复,意味着不发生求和。
  • 输出下标ij意味着将ij组合成一个二维输出。这本质上是将两个一维数组通过广播机制进行相乘。

广播机制在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, keepdims=True) # 计算加权和 # attn_weights 形状 (b, s, s), V 形状 (b, s, d_v) # 对键位置 k 进行求和,得到每个查询位置 q 对应的上下文向量 context = np.einsum(‘bqk,bkv->bqv’, attn_weights, V)

这段代码清晰地展示了信息流动:QK相互作用产生注意力权重,权重再作用于V。下标bqd,bkd->bqkbqk,bkv->bqv完美刻画了张量间的缩并关系,比用一堆transposematmul要直观太多。

4. 性能、优化与内存细节

虽然einsum在表达上非常简洁,但它的性能并非总是最优。理解其背后的机制有助于你在正确的地方使用它。

4.1 einsum的执行路径与优化

当你调用np.einsum时,NumPy 内部会经历以下几个步骤:

  1. 解析下标字符串:解析输入输出下标,确定哪些维度需要求和,以及结果的形状。
  2. 路径优化:对于涉及多个操作数(两个以上)的复杂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参数是关键。optimize=False禁用优化,按直观顺序计算,可能很慢。optimize=True(默认)会尝试寻找较优路径。optimize=‘optimal’会寻找理论最优路径,但对于大量操作数可能搜索较慢。optimize=‘greedy’是默认的启发式算法,在速度和效果间取得平衡。
  3. 执行计算:根据优化后的路径,调用底层(通常是BLAS库)的矩阵运算或执行循环进行求和。

4.2 与专用函数的性能对比

对于常见的、有专用函数的操作,直接调用专用函数通常更快,因为它们是高度优化的。

  • 矩阵乘法np.dot(A, B)np.matmul(A, B)通常比np.einsum(‘ij,jk->ik’, A, B)稍快,因为它们直接调用高度优化的BLAS库(如OpenBLAS, MKL)。einsum需要经过一层解析和分发。
  • 转置A.TA.transpose()是视图操作,几乎零成本。np.einsum(‘ij->ji’, A)会创建一个新的数组,有内存分配和复制开销。
  • np.trace(A)np.einsum(‘ii->’, A)性能接近,但专用函数可能略有优势。
  • 点积np.dot(a, b)np.einsum(‘i,i->’, a, b)性能类似。

那么,什么时候该用einsum

  1. 操作复杂,没有现成的专用函数:这是einsum的主场。比如前面提到的张量缩并、批量特定维度乘法等。
  2. 代码可读性优先:即使有替代方案,如果einsum表达式能一眼看清计算意图(如注意力机制),为了代码的清晰和可维护性,牺牲一点微不足道的性能是值得的。
  3. 原型设计和探索:在尝试新的数学公式或模型结构时,用einsum快速验证想法非常方便。

实操心得:路径优化是双刃剑。对于非常复杂的表达式(如涉及4个以上张量),开启optimize=True能带来数量级的性能提升。但优化过程本身有开销。对于在循环中反复调用的、非常简单的einsum(如简单的矩阵乘法),关闭优化optimize=False有时反而更快,因为避免了每次调用时的路径分析开销。我的经验法则是:在循环外预先计算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. 跨框架的einsum:NumPy, PyTorch, TensorFlow

einsum的概念已被主流深度学习框架广泛采纳,语法几乎完全一致,这带来了极大的便利。

  • NumPy:np.einsum(subscripts, *operands, out=None, dtype=None, order=‘K’, casting=‘safe’, optimize=False)
  • PyTorch:torch.einsum(equation, *operands)。PyTorch 的einsum支持自动微分,可以无缝嵌入神经网络中。在GPU上,它能调用优化的CUDA内核。
  • TensorFlow:tf.einsum(equation, *inputs, **kwargs)。同样支持GPU和自动微分。

框架间的重要差异:

  1. 优化策略:NumPy的optimize参数在PyTorch和TensorFlow中不一定有完全相同的实现或默认行为。PyTorch的einsum底层会尝试将操作映射到一系列基础的mm,bmm,sum等操作上。TensorFlow 的tf.einsum会尝试使用MatMul等核心操作。
  2. 广播规则:虽然都支持广播,但细微规则可能略有不同,在编写跨框架代码时需要注意。
  3. 性能:对于能在底层映射到高度优化算子(如torch.bmm,tf.matmul)的einsum表达式,框架的einsum性能可能接近专用函数。但对于非常特殊的缩并,可能退化为通用的、较慢的实现。
  4. 动态形状:在PyTorch和TensorFlow的图模式下(如tf.function, TorchScript),einsum对动态形状的支持可能不如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 ‘…’你使用了省略号...,但可能用法不对,或者输出下标指定的维度数与结果的实际维度数不符。

调试技巧:

  1. 画图:在纸上画出每个张量的方块图,用箭头连接需要求和(缩并)的维度。这能直观地检查维度是否对齐。
  2. 分步验证:对于复杂的表达式,拆分成多个简单的einsum或使用einsum_path查看中间步骤的形状。
  3. 使用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], dtype=np.int8) b = np.array([100, 200], dtype=np.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 最佳实践总结

  1. 从简单开始:先用einsum实现你熟悉的操作(如点积、矩阵乘),确保理解正确。
  2. 下标命名有意义:虽然字母是任意的,但使用有助记忆的字母(如batch,channel,height,width,input,output)能极大提升代码可读性。
  3. 优先使用专用函数:对于dot,matmul,trace,transpose等简单操作,直接调用专用函数通常更优。
  4. 复杂表达式先优化:对于涉及三个及以上操作数的运算,务必使用np.einsum_path检查优化路径,特别是当数据量较大时。
  5. 关注内存:留意路径优化报告中的“最大中间大小”,确保其不会超出可用内存。
  6. 测试与验证:用随机数据和小规模数据验证einsum表达式的结果是否正确,可以对比使用循环实现的“朴素”版本的结果。
  7. 文档化:在复杂的einsum表达式旁添加注释,说明每个下标字母的含义和运算的物理意义。

einsum是一个需要稍加练习才能熟练掌握的工具,但一旦掌握,它就会成为你处理多维数组运算的思维语言。它强迫你清晰地思考每个维度的去向,这种清晰性本身就能减少bug。下次当你面对一堆transposereshape感到头晕时,不妨试试用einsum来重新表述你的问题,很可能你会发现一条更清晰的道路。

http://www.jsqmd.com/news/1324181/

相关文章:

  • 从BERT到RAG再到Agent-First搜索:AI搜索引擎架构演进图谱(附23家厂商技术栈拆解与兼容性矩阵表)
  • 2025-AAAI《Anchor Learning with Potential Cluster Constraints for Multi-view Clustering》
  • 极值点、驻点与拐点:概念辨析与判别方法全解析
  • 基于Docker部署AI客户端API网关:打破AI应用孤岛
  • MyBatis注解开发实战:从基础CRUD到动态SQL与混合使用策略
  • C#实战:从递归算法到可视化交互的汉诺塔游戏开发指南
  • 基于RAG与本地大模型的轻量级智能文档问答系统实践
  • 双擎并驱·全链路并行:KFS让TB级异构增量同步秒级到达
  • 2026精选:景区指示牌源头工厂怎么选——安徽思扬标识工程有限公司深度剖析 - 装修教育财税推荐2026
  • 《辣知化智》13来自远古的意识密码
  • FPGA驱动LCD:从HD44780时序到Verilog状态机实战
  • 生态旅游教学实训实验室技术架构:VR资源库+3DGS+人景合一+数字人四层工具链
  • Mermaid+AI:用自然语言生成流程图,提升技术文档与设计效率
  • Spring事务失效的15种常见场景与解决方案
  • PCB批量阻抗校准体系搭建,三级联动修正方案
  • 2025届毕业生推荐的十大降重复率方案解析与推荐
  • Claude Code记忆系统:AI编程助手的上下文持久化与智能召回实战
  • 2026 年现阶段,苏尼特右旗热门的1085无缝钢管源头厂家选哪家,别再乱买理财了,这款5年108倍收益的无门槛产品,真的能稳赚吗? - 鉴选官
  • 避开这七个坑,你的网络安全自学之路能少走三年弯路
  • 回溯算法精解:从N皇后问题掌握递归、剪枝与状态搜索
  • 英雄联盟海斗模式录播学习法:从高手对局中系统提升游戏技术
  • 技术博客创作指南:从第一篇到持续发布
  • Sunshine游戏串流服务器:打造家庭游戏云的终极指南
  • 统计显著性:从A/B测试到数据驱动决策的核心原理与实践
  • 企业AI落地:超越模型选型,构建分层架构的实战指南
  • 介绍生物素化转铁蛋白Biotin-Transferrin,生物素-转铁蛋白Biotin-Tf的制备方法
  • Sunshine游戏串流完整指南:5步搭建你的私人游戏云平台
  • AI视频生成框架LibTV本地部署与实战指南:从环境搭建到工作流调优
  • 【万有无界技术解析】阿里多角色Agent协作工作台如何交付复杂项目
  • 2026 年现阶段渝中专业的GB9948石油裂化无缝钢管供应厂家哪家强,用对它,能让石油裂化设备寿命直接翻倍? - 行业推荐【认证官】