理解 NumPy 的 einsum¶
原文:Understanding Numpy's einsum 作者:Eli Bendersky 发布时间:2025-03-22 译文版本:v0.1
这是一篇关于如何使用 numpy.einsum 的简要说明与使用手册。einsum 允许我们使用爱因斯坦记号来计算多维数组上的运算。本文主要关注 einsum 的显式模式(在下标字符串中使用 -> 明确指定输出维度),以及机器学习论文中常见的用例;当然也会简要介绍其他模式。
基础用例:矩阵乘法¶
先从一个基础示例开始:使用 einsum 进行矩阵乘法。本文中,A 和 B 始终表示下面这两个矩阵:
>>> A = np.arange(6).reshape(2,3)
>>> A
array([[0, 1, 2],
[3, 4, 5]])
>>> B = np.arange(12).reshape(3,4)+1
>>> B
array([[ 1, 2, 3, 4],
[ 5, 6, 7, 8],
[ 9, 10, 11, 12]])
A 和 B 的形状允许我们计算 A @ B,得到一个 (2,4) 矩阵。使用 einsum 也可以完成同样的运算:
einsum 的第一个参数是下标字符串(subscript string),它描述了要对后续操作数执行的运算。其格式是:以逗号分隔的输入列表,后面跟着 -> 和一个输出。随后可以提供任意数量的位置操作数;这些操作数与下标中指定的、以逗号分隔的输入一一对应。对于每个输入,它的形状都表示为一串维度标记,例如 i(任意单个字母)。
在我们的例子中,ij 指代矩阵 A——用 (i,j) 表示它的形状;jk 指代矩阵 B——用 (j,k) 表示它的形状。虽然这些维度标记在下标中是符号性的,但当 einsum 接收到实际操作数并被调用时,它们就会变成具体的维度。这是因为此时操作数的形状已经确定。
下面是对 einsum 工作方式的一个简化心智模型(更完整的说明请阅读《einsum 的教学式实现》):
- 下标的输出部分指定了输出数组的形状,并使用输入维度标记来表达。
- 如果某个维度标记在输入中重复出现、但没有出现在输出中,那么它就会被约简(contract)(求和)。在我们的例子中,
j重复出现(且没有出现在输出中),因此它会被约简:每个输出元素[ik]都是第一个输入的第i行与第二个输入的第k列的点积。
只要调换输出的形状,我们就能轻松转置输出:
这等价于 (A @ B).T。
阅读机器学习论文时,我发现,即使对于基础矩阵乘法这样简单的情况,作者也经常更偏好使用 einsum 记号,而不是普通的 @(或者它的函数形式,如 np.dot 和 np.matmul)。这很可能是因为 einsum 的写法具有自文档化的特点,能帮助作者更明确地推理各个维度。
批量矩阵乘法¶
当输入的 ndim1 增加时,用 einsum 代替 @ 来充当矩阵乘法的文档说明,会显得更有意义。例如,我们可能希望在一次操作中对一整批输入执行矩阵乘法。假设我们有下面这些数组:
这里的 6 是批次维度(batch dimension)。我们要将一批 6 个 (2,3) 矩阵,与一批 6 个 (3,4) 矩阵相乘;Ab 中的每个矩阵都与 Bb 中对应的矩阵相乘。结果的形状是 (6,2,4)。
我们可以通过 Ab @ Bb 执行批量矩阵乘法——在 NumPy 中,这种写法直接就能工作:第一个数组的最后一个维度会与第二个数组的倒数第二个维度进行约简。这个操作会对最后两个维度之前的所有维度重复执行。输出形状是预期的 (6,2,4)。
使用 einsum 记号也可以完成同样的操作,而且写法更具自文档性:
这与 Ab @ Bb 等价,但下标字符串可以让我们用单个字母为维度命名,从而更容易理解正在发生什么。例如,在这里 b 可以代表批次(batch),m 和 n 可以代表序列长度,d 则可以代表某种模型维度或深度。
注意:虽然 b 在输入下标中重复出现,但它也出现在输出中;因此它不会被约简。
输出维度的排列顺序¶
einsum 下标中输出维度的顺序,不仅能让我们完成矩阵乘法;我们还可以用它转置任意维度:
在多维批量数组乘法中,这项能力经常与矩阵乘法结合使用,以精确指定维度的排列顺序。下面的例子直接取自 Noam Shazeer 的论文 Fast Transformer Decoding。
在批量多头注意力一节中,论文定义了下面这些数组:
M:形状为(b,m,d)的张量(批次、序列长度、模型深度)P_k:形状为(h,d,k)的张量(头的数量、模型深度、键的头大小)
我们先定义一些维度大小常量和随机数组:
>>> m = 4; d = 3; k = 6; h = 5; b = 10
>>> Pk = np.random.randn(h, d, k)
>>> M = np.random.randn(b, m, d)
论文使用一次 einsum 运算来计算所有键:
注意,这里既约简了 d 维度,也对输出进行了排列,使批次维度位于头维度之前。理论上,我们也可以通过下面的写法反转这两个维度的顺序:
实际上,输出可以采用任意顺序。显然,对于眼前这个具体操作,bhmk 才是合理的顺序。这里需要特别强调的是,与简单的 M @ Pk 相比,einsum 的写法具有更好的可读性;在后者中,参与运算的维度就不那么清楚了2。
对多个维度进行约简¶
一次 einsum 可以约简多个维度,下面是同一篇论文中的另一个例子:
>>> b = 10; n = 4; d = 3; v = 6; h = 5
>>> O = np.random.randn(b, h, n, v)
>>> Po = np.random.randn(h, d, v)
>>> np.einsum('bhnv,hdv->bnd', O, Po).shape
(10, 4, 3)
h 和 v 都同时出现在输入下标中,但没有出现在输出中。因此,这两个维度都会被约简——输出中的每个元素,都是沿 h 和 v 两个维度求和的结果。如果没有 einsum,要完成这个操作会麻烦得多!
转置输入¶
指定 einsum 的输入时,我们可以通过重新排列维度来转置输入。回想一下形状为 (2,3) 的矩阵 A;我们不能将 A 与自身相乘,因为形状不匹配,但可以像 A @ A.T 那样将它与自身的转置相乘。使用 einsum,可以这样写:
注意下标中第二个输入的维度顺序:这里是 kj,而不是前面使用的 jk。由于 j 仍然是输入中重复、但被省略在输出中的标记,因此它就是被约简的维度。
两个以上的参数¶
einsum 支持任意数量的输入;假设我们希望使用下面这个数组,将 A 和 B 进行链式矩阵乘法:
我们得到:
使用 einsum,可以这样写:
>>> np.einsum('ij,jk,kp->ip', A, B, C)
array([[ 900, 1010, 1120, 1230, 1340],
[2880, 3224, 3568, 3912, 4256]])
这里同样可以看到,明确命名各维度是很好的自文档化手段。
einsum 的教学式实现¶
上面介绍的 einsum 工作方式的简化心智模型并不完全准确,不过,它绝对足以帮助我们理解最常见的用例。
我读过很多网上介绍“einsum 如何工作”的文章,但遗憾的是,它们都存在类似的问题;委婉地说,至少它们都是不完整的。
我发现,实现一个基础版本的 einsum 很容易;而且,这个实现不仅能更好地解释 einsum 的工作方式,也能提供比其他尝试更好的心智模型3。那么我们就开始吧。
我们将基础矩阵乘法 'ij,jk->ik' 作为贯穿示例。
这个计算有两个输入,因此先编写一个接收两个参数的函数4:
下标中的标记指定了这些输入的维度,因此我们先显式定义各个维度的大小(如果某个标记在多个输入中重复出现,还要用 assert 确认这些大小相互兼容):
输出形状是 (i,k),所以我们可以创建一个空的输出数组:
然后为输出数组中的每个元素生成循环:
那么,这个循环中要放什么?现在该回头看看下标中的输入了。由于 j 标记发生了约简,这意味着要沿这个维度求和:
for i in range(i_size):
for k in range(k_size):
for j in range(j_size):
out[i, k] += __a[i, j] * __b[j, k]
return out
注意我们在循环体中如何访问 out、__a 和 __b;这直接来自下标 'ij,jk->ik'。事实上,einsum 正是这样从爱因斯坦记号发展而来的——后文会进一步介绍。
再看一个如何用这种方式推理 einsum 的例子,考虑“对多个维度进行约简”一节中的下标:
我们可以立即按照下标写出输出的赋值语句:
剩下的就是确定循环。前面已经讨论过,外层循环遍历输出维度,输入中被约简的维度则额外形成内层循环(本例中是 v 和 h)。因此,完整实现如下(省略 *_size 变量的赋值和维度检查):
for b in range(b_size):
for n in range(n_size):
for d in range(d_size):
for v in range(v_size):
for h in range(h_size):
out[b, n, d] += __a[b, h, n, v] * __b[h, d, v]
如果 einsum 下标中没有任何被约简的维度,会发生什么?这种情况下就没有求和循环;外层循环(为输出数组的每个元素赋值)只需将对应的输入元素相乘即可。下面是一个例子:'i,j->ij'。和前面一样,我们先设置维度大小和输出数组,然后遍历每个输出元素:
def calc(__a, __b):
i_size = __a.shape[0]
j_size = __b.shape[0]
out = np.zeros((i_size, j_size))
for i in range(i_size):
for j in range(j_size):
out[i, j] = __a[i] * __b[j]
return out
由于输入中没有任何不出现在输出中的维度,因此不存在求和。这个计算的结果就是两个一维输入数组的外积。
我已经在 GitHub 上放置了这套转换的完整且带有详细注释的实现。其中的 translate_einsum 函数接收一个 einsum 下标,并生成实现该下标的 Python 函数文本。
爱因斯坦记号¶
这种记号之所以以 Albert Einstein 命名,是因为他在 1916 年关于广义相对论的奠基性论文中将它引入了物理学。Einstein 需要用张量表达繁琐的嵌套求和,因此使用这种记号来简化表达。
在物理学中,张量通常同时具有下标和上标(分别表示协变分量和逆变分量),我们经常会遇到这样的方程组:
我们可以使用变量 i 将它们合并成一个求和:
接着注意到,j 在求和中重复出现(一次位于下标,一次位于上标),所以可以把它写成:
这里隐含了求和;这就是爱因斯坦记号的核心。
补充知识:爱因斯坦求和约定的三条规则¶
MathWorld 将爱因斯坦求和记号概括为三条规则:
- 重复出现的指标默认进行求和。 例如,
j在a_{ij}A^j中出现两次,因此表示对j求和。 - 同一个指标在一个项中最多出现两次。 如果某个指标在同一个项中出现三次,就不符合这套约定。
- 每一项必须包含相同的非重复指标。 非重复指标是自由指标,它们决定表达式的输出分量;不同项中的自由指标必须一致。
在物理学中,这套约定通常还会与 Kronecker delta 和置换符号一起使用,并且可以自然地处理协变张量与逆变张量的上标、下标。详见MathWorld 对 Einstein Summation 的介绍。
对于 NumPy einsum,这些规则提供了很有用的直觉,但两者并不完全相同:einsum 下标中的字母通常只是维度标签,不会自动表达物理学中的协变或逆变含义;重复且未出现在输出中的标签则表示要约简(求和)的维度。
细心的读者可能会注意到,原来的方程组很容易表示成矩阵-向量乘法,但还需要记住两点:
- 矩阵记号是在 Einstein 研究广义相对论之后才开始流行于物理学中的(事实上,最早引入它的是 Werner Heisenberg,时间是 1925 年)。
- 爱因斯坦记号可以扩展到任意数量的维度。矩阵记号在二维情况下很有用,但在更高维度中就很难可视化和操作。在二维情况下,矩阵记号与 Einstein 的记号是等价的。
应该不难看出,这种记号与本文讨论的 einsum 下标之间的对应关系。概念上,einsum 的隐式模式与爱因斯坦记号更加接近。
隐式模式的 einsum¶
在 einsum 的隐式模式中,输出规格(-> 以及其后的标记)不存在。相反,输出形状是根据输入标记推断出来的。例如,下面是二维矩阵乘法:
在隐式模式中,每个输入内标记的字典序很重要,因为它决定了输出维度的顺序。例如,如果我们希望得到 (A @ B).T,可以这样写:
由于 h 在字典序中排在 i 之前,这等价于显式下标 'ij,jh->hi';而原始的隐式矩阵乘法下标则等价于 'ih,jk->ik'。
据我所知,隐式模式在机器学习代码和论文中并不常用。在我看来,与显式模式相比,它牺牲了很多可读性,却几乎没有因为省略输出标记而获得什么好处。
-
这里的 ndim 是指
numpy.ndim,也就是数组中的维度数量。它有时也被称为阶(rank),但这容易造成混淆,因为rank 已经是线性代数中另一个概念的名称。 ↩ -
我个人认为,NumPy 及其衍生库(如 JAX、PyTorch 和 TensorFlow)最大的缺点之一,是没有办法为运算的形状添加注解并进行检查。这使得很多代码的可读性不如本来可以达到的程度。在某种程度上,
einsum缓解了这个问题。 ↩ -
最早见于 Stack Overflow 的这个回答。 ↩
-
这里使用下划线的原因,是为了避免与可能名为
a和b的维度标记发生冲突。由于我们正在生成代码,变量遮蔽是一个常见问题;如果想了解更多有趣内容,可以参阅卫生宏(hygienic macros)。 ↩