在数值计算和深度学习领域,经常能碰到一堆数字摞在一起的情况,比如矩阵、三维数组甚至更高维的数据。要把这些数据按照特定规则“揉捏”成新的形状,如果只会用最笨的循环,代码不仅写得像裹脚布,还慢得让人抓狂。今天要聊的主角 np.einsum 就是专门干这个的,它能把复杂到看一眼就头晕的张量运算,写成一行清清楚楚的“下标注”。这东西一旦用顺手,你会觉得自己之前的很多代码都是在“搬砖”而不是在写程序。

一、从一次“搬砖”说起:什么是张量收缩

先别被“张量收缩”这个专业词吓到。咱们把张量简单理解成“装着数字的多维盒子”。一个值,是零维盒子;一队数,是一维盒子(向量);一个表格,是二维盒子(矩阵);一堆表格摞起来,就是三维盒子(张量)。现实中的张量往往不止三维,但理解方式一样。

“收缩”是什么意思?通俗讲,就是把两个盒子上的某些“轴”配对,然后按规则相乘并累加。比如两个矩阵相乘,本质是左边矩阵的行和右边矩阵的列做一对一的加权求和。这种操作在物理、机器学习、图像处理里到处都是,但手动实现时常常要写好几层循环,别人读起来累,你自己也容易错。

np.einsum 的做法是把你打算“怎么配对轴”的规则直接写出来,其他的循环全部交给底层的优化代码去跑。它不关心你的张量是几维,只要规则说清楚,它就能高效完成运算。

1.1 张量收缩的生活化比喻

想象你手里有一叠购物小票,每张小票上写着“商品名、单价、数量”。你想算总花费,脑子里做的其实就是一次收缩:把“单价”和“数量”两个轴配对相乘,然后把所有商品的结果加起来。这个过程中,商品名、单价、数量都是“轴”,你挑出其中两个“轴”做乘法,再消灭掉一个“轴”(求和),剩下的就是总钱数。np.einsum 干的就是这种活,只不过它处理的是数字数组,而且速度飞快。

二、np.einsum 的基本玩法

np.einsum 的全称是“爱因斯坦求和约定”的实现。它的核心思想是:用一个简单的字符串描述张量各维度之间的对应关系。比如 'ij,jk->ik' 就表示“下标为 i、j 的二维数组”与“下标为 j、k 的二维数组”进行运算,最终输出“下标为 i、k”的二维数组。箭头左边是输入张量的下标,箭头右边是输出张量的下标。如果下标在箭头左边出现了,但没出现在右边,就说明它被“收缩”(求和)掉了。

看一个最简单的例子,计算两个一维数组的点积:


# 技术栈:Python + NumPy
import numpy as np

a = np.array([1, 2, 3])
b = np.array([4, 5, 6])

# 'i,i->' 的意思是:两个数组的下标都是 i,输出没有下标
# 没有下标意味着输出是一个标量(数)
# 等价于 sum(a * b)
result = np.einsum('i,i->', a, b)

print(result)  # 输出:32

这种写法的好处是,运算规则一眼就能看出来。不用去查 NumPy 函数名,也不用去猜 np.dotnp.matmul 谁是谁的别名。

2.1 基本语法拆解

规则字符串的格式是:输入部分用逗号分隔每个张量的下标,箭头后面是输出下标。如果省略箭头和输出部分,哪个下标在输入里只出现一次,就保留;哪个下标在输入里出现了多次(比如 i 在第一个数组和第二个数组里都有),就被当做求和轴。不过为了代码清晰,强烈建议永远把箭头和输出写全,别偷懒。下面是省略箭头的等价写法:


# 技术栈:Python + NumPy
import numpy as np

x = np.array([[1, 2], [3, 4]])
y = np.array([[5, 6], [7, 8]])

# 省略箭头时,'ij,jk' 会自动补成 'ij,jk->ik'
# 因为 i 和 k 各出现一次可以保留,j 出现了两次会被求和
result = np.einsum('ij,jk', x, y)

print(result)
# 输出:
# [[19 22]
#  [43 50]]

这个结果就是一个标准的矩阵乘法。如果你平时用 np.matmul,结果完全一样,但 einsum 的表达方式更接近“数学公式”。

三、用生活场景理解下标规则

很多人第一次接触 einsum 时,觉得下标符号太抽象。其实你只需要记住一条规则:相同的下标就是“要配对”的轴,输出里没出现的下标就是“要消失”的轴。下面用几个超常见的场景来加深感觉。

3.1 矩阵乘法

矩阵乘法是“行乘列”的收缩运算。假设矩阵 A 的形状是 (m, n),矩阵 B 的形状是 (n, p),结果矩阵 C 的形状是 (m, p)。下标规则就是 'ij,jk->ik'。其中 j 是内部维度,被求和干掉了。


# 技术栈:Python + NumPy
import numpy as np

A = np.array([[1, 2],
              [3, 4],
              [5, 6]])  # 形状 (3, 2)

B = np.array([[7, 8, 9],
              [10, 11, 12]])  # 形状 (2, 3)

# 用 einsum 做矩阵乘法
C = np.einsum('ij,jk->ik', A, B)

print(C)
# 输出:
# [[ 27  30  33]
#  [ 61  68  75]
#  [ 95 106 117]]

如果不用 einsum,你可能得写三层循环或调用 np.matmul。但 np.matmul 只能做“标准”的矩阵乘法,一旦你要做“只交换某些轴”的运算,它就绕了。

3.2 批量矩阵乘法

在处理一批图像或一批句子时,经常要同时对很多个小矩阵做乘法。假设你有两批矩阵,每批都有 batch 个,形状分别是 (batch, m, n) 和 (batch, n, p),要得到 (batch, m, p) 的批量结果。einsum 只需要在规则前面加一个 batch 下标就行。


# 技术栈:Python + NumPy
import numpy as np

# 模拟一个小批量:2 个样本,每个样本是 3x2 矩阵
batch_a = np.random.rand(2, 3, 2)
# 模拟另一个小批量:2 个样本,每个样本是 2x4 矩阵
batch_b = np.random.rand(2, 2, 4)

# 下标 'bij,bjk->bik':b 是批次轴,不参与收缩
output = np.einsum('bij,bjk->bik', batch_a, batch_b)

print("批量结果形状:", output.shape)
# 输出:批量结果形状: (2, 3, 4)

这一个规则就把“对每一批分别做矩阵乘法”的循环和内部逻辑全部封装了,代码非常干净。

3.3 矩阵转置、求和、对角线等小操作

einsum 不仅能做乘法,还能做很多“单张量”操作。比如转置就是交换下标顺序,求对角线就是让两个下标相等,求和就是所有下标都不写在输出里。


# 技术栈:Python + NumPy
import numpy as np

M = np.array([[1, 2, 3],
              [4, 5, 6],
              [7, 8, 9]])

# 转置:'ij->ji'
transposed = np.einsum('ij->ji', M)
print("转置结果:")
print(transposed)
# 输出:
# [[1 4 7]
#  [2 5 8]
#  [3 6 9]]

# 求所有元素之和:'ij->'
total = np.einsum('ij->', M)
print("所有元素之和:", total)  # 输出 45

# 提取对角线:'ii->i'(但注意这里 M 的对角线是 1、5、9)
diag = np.einsum('ii->i', M)
print("对角线元素:", diag)
# 输出:对角线元素: [1 5 9]

看,一个函数干了 transposesumdiagonal 三个函数的活。而且写法直观,不用记新 API。

四、真实案例:注意力机制中的分数计算

在自然语言处理里,Transformer 的注意力机制核心就是一个张量收缩。假设我们有查询矩阵 Q、键矩阵 K 和值矩阵 V,形状都是 (batch, seq_len, dim)。计算注意力分数时,需要把 Q 和 K 的最后两维做点积,得到 (batch, seq_len_q, seq_len_k) 的分数矩阵。用 einsum 一行就能搞定。


# 技术栈:Python + NumPy
import numpy as np

# 模拟一个小 batch:2 个句子,每个句子 3 个词,每个词用 4 维向量表示
batch = 2
seq_len_q = 3
seq_len_k = 3
dim = 4

# 随机生成 Q、K、V
Q = np.random.rand(batch, seq_len_q, dim)
K = np.random.rand(batch, seq_len_k, dim)
V = np.random.rand(batch, seq_len_k, dim)

# 计算 Q 和 K 的点积分数
# 下标 'bqd,bkd->bqk' 中,b 是批次轴,q 和 k 是序列位置
# d 是特征维度,被求和干掉,最后得到每个 (q, k) 位置的相似度分数
scores = np.einsum('bqd,bkd->bqk', Q, K)

print("注意力分数形状:", scores.shape)
# 输出:注意力分数形状: (2, 3, 3)

# 接下来对分数做 softmax,再与 V 加权求和
# 先把分数转为概率(这里简单演示,实际需要沿第 2 个轴做 softmax)
weights = np.exp(scores) / np.sum(np.exp(scores), axis=-1, keepdims=True)

# 用 einsum 做加权求和新:'bqk,bkd->bqd'
# 这里 k 是序列位置轴,被求和干掉,也就完成了“把 V 的信息按注意力权重聚合起来”
context = np.einsum('bqk,bkd->bqd', weights, V)

print("上下文向量形状:", context.shape)
# 输出:上下文向量形状: (2, 3, 4)

这个例子完整展示了 einsum 在深度学习中的典型用法。如果不借助它,你需要用 np.matmul 配合 expand_dimstranspose 等做一堆轴操作,很容易在维度上翻车,而代码的可读性也差得多。

五、应用场景汇总

einsum 的应用范围非常宽,只要是涉及多维数组的“轴配对”运算,它都能派上用场。下面举几个具体场景:

  • 矩阵乘法系列:单矩阵乘法、批量矩阵乘法、向量的内外积。只需要改变下标,比如 'i,i->' 是点积,'i,j->ij' 是外积。
  • 高阶张量运算:在物理模拟、连续介质力学中,经常有类似 C[a,b,c,d] = A[a,b,e] * B[e,c,d] 的运算,einsum 可以直接写 'abe,ecd->abcd'
  • 数据降维与汇总:对特定维度求和、求均值、求对角线,或者把某些维度合并成对角矩阵。比如 'ijj->i' 可以提取三维张量中每个矩阵的迹。
  • 深度学习模型中的注意力机制:如上示例,计算 query 和 key 的得分,以及加权求和。
  • 量子计算模拟:量子态的张量网络收缩是典型的高维收缩场景,einsum 是常用的底层工具。
  • 图像处理:比如图像的通道变换、像素之间的局部加权求和,都可以抽象成张量运算。

六、技术优缺点

任何技术都有两面性,einsum 也不例外。下面聊聊它的好与坏。

6.1 优点

  • 代码极简、可读性强:一个规则字符串就能描述复杂的数学运算,比多层循环容易理解得多。
  • 减少维度错误:不用手动 reshape 和 transpose,只要下标正确,输出形状就符合预期。
  • 性能不错einsum 底层会调用优化的 BLAS 库,在能合并计算的时候会自动做优化,通常快于手写循环。
  • 表达力和通用性高:一个函数能代替 np.dotnp.matmulnp.sumnp.diagonalnp.transpose 等众多操作,而且可以自由组合。

6.2 缺点

  • 下标记忆有门槛:刚接触时容易搞不清 -> 左右该写啥,特别是高维场景,容易下标写错导致运行时错误或静默算错。
  • 规则字符串调试困难:如果结果不对,很难从字符串本身定位是哪个轴搞错了,需要自己拿小例子验证。
  • 性能并不总是最优:在某些特定形状下,手动使用专门的矩阵乘法函数(比如 np.matmul)可能更快,因为 einsum 需要解析字符串并尝试匹配最优路径,这个过程有一定开销。
  • 不够灵活:如果你想在运算过程中插入自定义操作(比如 clip 或除法),用 einsum 就不方便了,得拆成几步。

七、注意事项

einsum 时容易踩的坑,这里单独列出来提醒一下。

7.1 下标合法性

下标字母必须与张量的维度数一一对应。比如一个形状是 (3, 4) 的数组,下标必须恰好是两个字母,不能多也不能少。另外,同一个字母多次出现时,对应的维度长度必须相等,否则会报错。


# 技术栈:Python + NumPy
import numpy as np

a = np.random.rand(3, 4)
b = np.random.rand(4, 5)

# 下面这行是合法的:'ij,jk->ik',j 对应的维度长度都是 4
c = np.einsum('ij,jk->ik', a, b)

# 如果写成这样,就会报错:
# a 的第二个维度是 4,b 的第一个维度是 5,字母 j 两次对应的长度不一致
# np.einsum('ij,jk->ik', a, np.random.rand(5, 6))  # 运行时报维度不匹配错误

print(c.shape)  # 输出: (3, 5)

7.2 性能与内存

虽然在很多情况下 einsum 很快,但不要迷信它。当你的张量非常大,而且操作可以拆分成多个简单矩阵乘法时,分步使用 np.matmul 有时更快。另外,einsum 会创建中间结果,如果在循环里反复调用而且张量巨大,内存开销需要留意。一个实用的建议是:先用小数据验证下标正确,再放到真正的数据上跑。一旦结果不对,直接用 np.testing.assert_allclose 对照一个笨办法的循环实现来排查。

八、文章总结

np.einsum 是一个把“复杂到不想写循环”的张量收缩问题,压缩成一行字符串的利器。它的核心思想就一句话:相同的下标配对,输出里没出现的下标求和。掌握了这个规则,你就能用统一的视角看待矩阵乘法、转置、求和、批量运算和注意力机制里的各种操作。虽然它有一点学习门槛,但一旦习惯,你会发现自己写出来的代码不仅短,而且更容易让别人看懂。在科学计算和深度学习领域,这绝对是值得花点时间掌握的技巧。别怕下标,多拿几个小例子练练手,很快你就能体会到“一行代码搞定一切”的清爽感。