我最早被 einsum 圈粉,是在一次给 Transformer 做性能优化的时候。当时要重写一个多头注意力模块,手写 Q @ K.transpose(1,2) / sqrt(d_k) 再 @ V,然后做 permute、reshape,代码又臭又长。同事丢过来一行 einsum,把整个 attention 的批量矩阵运算全收缩进了下标表达式里,我当时就觉得这工具不一般。后来在 NumPy、PyTorch、JAX 里越用越多,算协方差、做张量分解、写多模态特征融合,几乎脏活累活都靠它简化。可以说,einsum(Einstein Summation)是一个从物理学家手稿里走出来、却特别适合做工程的高效张量运算工具,掌握了它,你的张量运算代码会简洁很多,而且往往比手动实现更快更省内存。
这篇文章不是概念科普,我会从直觉理解讲到底层优化,再给出可以直接抄的写法、性能实测和一些踩坑记录。适合正在做深度学习、科学计算、数据分析的工程师和数据科学家,也适合刚接触张量运算但想少走弯路的同学。最后给你的建议是:一天之内,把 einsum 用熟,不亏。
1. einsum 到底在表达什么
1.1 一次看懂爱因斯坦求和约定
这个名字容易吓退人,但背后的思想极其简单。爱因斯坦当年写广义相对论时,嫌求和符号写起来烦,干脆定了一条规则:在一项乘积里,如果同一个指标出现两次,就默认对这个指标求和。比如矩阵乘法 (C_{ij} = \sum_k A_{ik} B_{kj}),他直接写成 (C_{ij} = A_{ik} B_{kj}),不写求和符号。
einsum 就是把这条规则搬到了张量运算里。你用下标字符串描述“输入张量的每个维度叫什么名字”,然后在箭头右侧写下“输出张量保留哪些名字”,函数会自动帮你处理下标匹配、求和和广播逻辑。换句话说,你只需要表达计算意图,中间那些累加、转置、扩维的动作全部交给库去安排。
举个例子:两个二维矩阵做矩阵乘法,np.einsum('ik,kj->ij', A, B)。左边 ik 表示 A 的行叫 i 列叫 k,kj 表示 B 的行叫 k 列叫 j。右侧 ij 说明我们要的是 i 和 j。k 在左右只有一边出现?不对,k 在左边出现两次,它就是爱因斯坦求和约定里的哑指标,自动求和。所以这一句等价于 A @ B。
理解了这一点,后面所有复杂表达式都是同一个思路的扩展:重复下标求和,单次下标保留。
1.2 用一张表看清常见运算
我整理过一张速查表,后来一直贴在项目文档里,也分享给团队新人。很多看起来很绕的运算,无非是这几种模式的组合。
| 运算 | 手动实现 | einsum 写法 |
|---|---|---|
| 矩阵乘法 | A @ B |
'ik,kj->ij' |
| 批量矩阵乘法 | torch.bmm(A, B) |
'bik,bkj->bij' |
| 点积 | sum(a * b) |
'i,i->' |
| 外积 | a[:, None] * b[None, :] |
'i,j->ij' |
| 转置 | A.T |
'ij->ji' |
| 对角线 | np.diag(A) |
'ii->i' |
| 迹 | np.trace(A) |
'ii->' |
| 逐元素乘 | A * B |
'ij,ij->ij' |
| 降维求和 | A.sum(axis=1) |
'ij->i' |
| 批量点积 | sum(a * b, dim=1) |
'bi,bi->b' |
这张表最值得注意的一点是,所有运算都只用“下标名字”和箭头来描述。你不用管转置、广播、扩维的顺序,库会根据下标关系自动编排。这极大降低了写错维度的概率。
我在实际项目中见过很多新手花大量时间调 reshape 和 permute,最后发现是维度顺序搞反了。用 einsum 基本可以从头避免这类问题,因为下标不匹配时,函数会直接报错,而不会给你一个形状对但语义全错的张量。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 为什么 einsum 能这么快:底层优化逻辑
2.1 省掉中间张量的艺术
很多人一开始觉得 einsum 就是语法糖,包了一层循环或一堆 reshape。其实它背后的优化能力远超这个判断。关键在于,einsum 描述的是整个计算的“收缩路径”,后端可以在高层统一做调度,而不是机械地一步步执行你手写的顺序。
举个具体例子:计算多个矩阵的乘积 A @ B @ C @ D,手动写法可能是 ((A @ B) @ C) @ D,每一步都会生成一个中间张量。中间张量一多,内存占用和访存开销都会变大。而 einsum 的下标表达式 'ij,jk,kl,lm->im' 可以让后端在计算前分析:哪些中间结果可以合并?哪些维度应该先收缩?选哪条路径最省算力和内存?这跟编译器做表达式重组是一个道理。
我实测过一个小矩阵链乘法:四组 64×64 矩阵连续相乘,在 NumPy 里用 einsum 有时比手动连乘快 20%~40%。如果矩阵更大,这个收益还会更明显。它不是玄学,是减少中间张量分配和复制带来的真实收益。
2.2 后端如何选择最优收缩路径
像 opt_einsum 这类库,会把 einsum 表达式转换成一个图,然后尝试不同的“配对收缩顺序”,估算每种的浮点运算量和临时内存,选一个最优解。PyTorch 从某个版本开始也接入了类似的优化逻辑,JAX 的 XLA 编译器同样会对 einsum 做布局优化和融合。
对我们使用者来说,不需要记这些细节,但要知道一个核心点:einsum 的加速上限,取决于后端到底有没有做路径优化。在 NumPy 的老版本里,einsum 其实实现得很保守,速度不一定比得上 BLAS 加速后的 @。到了 PyTorch、JAX 或新版 NumPy 里,优化就要激进得多。
因此搭配建议是:如果只是简单矩阵乘法,直接用 @ 或 torch.matmul 就行,BLAS 已经把性能压榨得够强。einsum 的优势场景是复杂一点的收缩运算,比如批量点积、多张量收缩、混合转置和求和,这些场景下你手写的组合很难超过后端的全局优化。
2.3 为什么 einsum 能和公式直接对应
深度学习论文里经常出现像 (H = (Q K^\top) V) 这样的记号。你照着写代码时,可能需要把维度展开、加 batch 维度、注意 transpose 的轴顺序,中间很容易出错。
einsum 的价值在于它保留了和数学公式几乎一致的表达层级。'bqhd,bkhd,bkhd->bqhd' 这类写法,别人看一眼下标就能理解你在算什么,代码和公式相互印证。这种“可读即推导”的特性,在项目维护和代码 review 里特别实用,因为后续接手的人不用一层层追踪 permute 后张量每个维度到底是什么意思。
3. 一字一句搞懂 einsum 语法
3.1 输入输出下标与箭头规则
einsum 的标准语法是 einsum(subscripts, *operands)。subscripts 由三部分组成:输入下标、箭头、输出下标。没有箭头时,默认输出为所有没有重复出现的下标按字母表排序,这个叫隐式模式。比如 np.einsum('ij,jk', A, B) 等价于 np.einsum('ij,jk->ik', A, B)。
我建议任何时候都写显式箭头,不要依赖隐式模式。显式模式不仅语义更清楚,还能避免某些奇怪的默认排序导致的结果不符合预期。项目里写 'ij,jk->ik' 比写 'ij,jk' 多不了几个字符,但后续读代码的人会轻松很多。
输入下标里的字母,同一个张量各维度必须用不同字母,不同张量可以用同一个字母表示“这个维度要发生交互”。输出下标必须是输入下标中出现过的字母,且必须按你想要的维度顺序排列。输出里没出现的输入下标,意味着该维度会被求和(如果出现两次)或被丢弃(输出里没有就相当于规约)。
3.2 省略号是用来处理批量维度的
真实数据通常不止二维。图片是 (B, H, W, C),序列是 (B, L, D),Transformer 的权重是 (D_out, D_in)。如果每次都用固定的 i、j 字母去匹配前面的批量维度,写起来会非常痛苦。
einsum 用省略号 ... 表示“这里有一批维度我懒得列了,自动匹配”。比如 np.einsum('...ij,...jk->...ik', A, B),就可以让批量矩阵乘法自动适配任意维度前缀,只要 A 和 B 在 ... 上的形状一致。
需要注意,省略号在同一个表达式里可以出现在多个输入中,但表示的前缀维度数量需要一致,否则会报错。举个错误例子:A 的形状是 (2, 3, 4, 5),B 的形状是 (3, 4, 5),写 '...ij,...jk->...ik' 时,A 的 ... 匹配 (2, 3) 还是别的,B 的 ... 匹配 (3,),两者数量不一致,运行时就会因为维度对不上而报错。这种问题需要先搞清楚想要的计算语义,把 B reshape 到 (1, 3, 4, 5) 再计算。
3.3 维度复用、广播和换轴的隐性操作
很多人在写 einsum 时会忽略“维度复用”这个能力。比如 'i,i->i' 表示逐元素相乘,'i,j->ij' 表示外积。这里 i 和 j 分别只出现一次,箭头右侧同时出现了 i 和 j,就相当于自动执行了 broadcast。
更皮的一点是,同一个输入张量里可以重复使用相同下标,比如 'ii->i' 取对角线,'ii->' 取迹。这种用法是从张量自身拿同一维度去运作,理解成“等价于一个循环在逐个元素取值再累加”就行。
还有一类隐藏能力:下标顺序就是输出维度顺序。比如 'bij->bji' 等于维度交换。这能让你在做后续运算时少写很多 .permute(1, 0)。
自己写着玩的时候,可以多试试让同一个下标出现在不同的位置,观察输出形状变化,很快就能建立起直觉。
4. 实战场景:一篇讲透多个高价值用法
4.1 用 einsum 重写多头注意力机制
这是 einsum 在深度学习中最能体现优势的场景之一。标准的多头注意力大概长这样:
- 输入
x形状(batch, seq_len, d_model) - 通过权重矩阵映射成 Q、K、V,都 reshape 成
(batch, seq_len, num_heads, head_dim) - 计算缩放点积注意力分数:
(Q @ K.transpose(-2, -1)) / sqrt(head_dim) - softmax 后再和 V 相乘
手写 QK^T 那一步,很多人会先把 Q 和 K 做 transpose,再用 torch.matmul,中间还要担心 batch 维度一致性。用 einsum 直接一行:
python复制scores = torch.einsum("b q h d, b k h d -> b h q k", q, k) / math.sqrt(head_dim)
注意这里把输出下标写成了 "b h q k",这样算出来的 scores 直接就是 (batch, heads, query_len, key_len) 的形状,省掉一次 transpose。后面和 V 做乘法时:
python复制out = torch.einsum("b h q k, b k h d -> b q h d", probs, v)
然后 reshape 回到 (batch, seq_len, d_model) 就行。我试过用 einsum 重写后,整个 attention 模块的代码行数减少了一半,而且形状语义清楚得多。更关键的是,在 PyTorch 里,这样写和显式 @ + transpose 的性能基本持平,不会像某些自以为是的写法那样带来额外开销。
4.2 批量计算协方差矩阵与统计量
数据分析里经常要算“每组样本的协方差矩阵”。比如有一批特征矩阵 X 形状 (N, D),要得到 X 经过中心化后的 X.T @ X / N。手写代码要两步:先减均值,再做转置矩阵乘。用 einsum 可以把“中心化后求和”和“二次型”压缩在一个表达式里吗?严格说,中心化需要显式先做,但后面的二阶矩计算可以直接用 einsum:
python复制# X_centered 形状 (N, D)
cov = torch.einsum("ni,nj->ij", X_centered, X_centered) / (N - 1)
这个表达式一秒内就能算出 D×D 协方差。如果是多维批量数据,比如 X 形状 (B, N, D),需要按 batch 分别算:
python复制cov_batch = torch.einsum("bni,bnj->bij", X_centered, X_centered) / (N - 1)
同样的套路,也可以做批量欧氏距离矩阵、批量 Gram 矩阵、批量白化操作。这些在 metric learning、特征分布分析里太常用了。
4.3 张量分解与多因子收缩
做推荐系统或信号处理时,经常会遇到 Tucker 分解、CP 分解这类模型。这些模型的核心运算就是多张量沿某些维度收缩。如果没有 einsum,你要写一串循环或反复 squeeze、matmul,复杂度极高。
举例:给定一个三阶张量 T 形状 (I, J, K),和一个因子向量 a 形状 (I,),想计算 (u = \sum_i a_i T_{ijk}) 的二维结果。einsum 一行:
python复制u = np.einsum("i,ijk->jk", a, T)
如果再来一个 b 向量,计算 (v = \sum_{ij} a_i b_j T_{ijk}):
python复制v = np.einsum("i,j,ijk->k", a, b, T)
这就是典型的多因子收缩。einsum 可以同时接收多个操作数,并根据同一个表达式把它们组合在一起。这在手动实现时你至少得拆成两步,还会产生中间张量。einsum 的路径优化在这里很值钱,因为它是直接在多个张量间寻找最佳收缩方案。
4.4 图像与点云处理中的批量变换
图像增广或点云变换时,经常需要对每个样本施加同一个可学习矩阵。比如点坐标 pts 形状 (B, N, 3),旋转矩阵 R 形状 (B, 3, 3),想得到每个点旋转后的坐标:
python复制rotated = torch.einsum("bij,bnj->bni", R, pts)
这里 bni 的输出下标使得旋转后的每个点仍是 (N, 3) 排列,不用事后做维度重排。类似的,如果是同一旋转矩阵作用到所有样本,因子变成 (3, 3):
python复制rotated = torch.einsum("ij,bnj->bni", R, pts)
这个写法在很多三维视觉项目里非常实用。我见过的不少点云代码,为了做这一步会先 tile 变换矩阵再 bmm,完全没有必要。einsum 能感受到矩阵维度大小,自动完成广播,省内存也省代码。
5. 性能实测:什么时候该用,什么时候别硬用
5.1 我做的几组基准对比
为了搞清楚 einsum 的实际性能,我在一台 GPU 服务器上做过几组测试。测试环境大概是 PyTorch 2.x + CUDA,测试了三种典型操作:矩阵乘法、批量点积、混合收缩。对比对象是“手动实现”和“einsum 实现”。
第一组:批量矩阵乘法 (B=128, M=64, K=128, N=64)。torch.bmm 和 torch.einsum("bik,bkj->bij") 的耗时几乎一样,einsum 甚至略快一点点,但差异在噪声范围内。
第二组:批量点积,形状 (B=1024, D=512)。手写 (a * b).sum(dim=1) 和 einsum("bd,bd->b") 也是基本打平。CPU 上,PyTorch 对 sum 有专门的 SIMD 路径,einsum 不一定能占便宜。
第三组:三个张量收缩,比如计算 a 形状 (B, D)、b 形状 (D, E)、c 形状 (B, E) 的某种组合。这时手写代码要么生成中间张量,要么做循环,einsum 有明确优势。我实测的一个场景中,einsum 比中间张量法快了约 15%,同时把峰值内存降低了一个层级。
表格里大概是这样:
| 运算类型 | 张量规模 | 手写耗时 | einsum 耗时 | 结论 |
|---|---|---|---|---|
| 批量矩阵乘 | (128,64,128,64) | 2.1 ms | 2.0 ms | 基本持平 |
| 批量点积 | (1024,512) | 0.5 ms | 0.5 ms | 持平 |
| 三张量收缩 | (128,256,128,256,128) | 8.6 ms | 7.4 ms | einsum 胜出 |
| 嵌套循环收缩 | 小规模多次 | 3.2 ms | 9.8 ms | 手写反而快 |
最后一行揭示了一个容易忽略的点:如果操作本身很小,比如只有几十个元素、循环简单,einsum 的函数调用和表达式解析开销可能反而超过手写循环。这时候不要为了“炫技”而强行 einsum。
5.2 einsum 不是银弹,这些情况建议手写
第一,操作过于简单且频繁出现时,比如单个向量点积,np.dot 或 sum(a*b) 就够了。einsum 理论上能表达,但并不总是快。
第二,和 FlashAttention 这类融合了特殊存储布局的专用算子相比,einsum 只是表达层工具,不具备按照 fused kernel 方式重排计算的能力。要榨干注意力性能,仍然需要专用 kernel,而不是 einsum。
第三,一些非收缩性的逐元素操作,比如 A + B、A * 2 + B,einsum 写起来别扭,性能也没有优势。直接用运算符。
第四,代码需要支持老版本库或兼容多种后端时,einsum 的行为差异可能成为隐藏坑。稍后会在常见问题里展开。
5.3 读懂一张性能火焰图的必要性
如果你打算在大型项目里全面铺开 einsum,我建议先跑一下 profiling,不要只看一两个 benchmark。einsum 的高层优化在某个表达式上可能非常优秀,但换一个表达式、换一种 shape、换一个后端,表现可能完全不同。
一个简单做法是,把现有关键路径里的核心运算分别用手写和 einsum 各实现一遍,放到同样的数据规模下,用 torch.profiler 或 line_profiler 比较。不要凭印象下结论,也不要只测小规模数据。之前我有一个同事坚持认为 einsum 一定优于 permute + matmul,结果在某个特定 shape 上完全相反,因为那个 shape 恰好能让矩阵乘走进特别优化的 BLAS 路径。所以实践法则很朴素:以你真实场景的 shape 为基准,跑一遍再决定。
6. 常见错误与调试心得:从入门到躺平
6.1 五种高频报错及解法
einsum 报错通常都和信息有关,但信息有时候不太直观。我把常见错误汇总成了一份速查表:
| 报错现象 | 可能原因 | 解决方案 |
|---|---|---|
| 输出下标但没在输入下标里出现 | 写错了输出维度名 | 检查箭头右侧字母是否都出现在左侧 |
| 维度数量对不上 | 操作数形状和下标不匹配 | 打印每个张量的 shape,反向核对字母代表的轴长度 |
| 省略号维度不一致 | 两个输入的前缀维度长度不一致 | 手动 reshape 统一前缀维度 |
| 类型被推断成 float64 或内存爆涨 | 大整数求和溢出或精度被抬高 | 显式转换 dtype |
结果形状是空元组 () |
所有下标都收缩了 | 这是合法的标量输出,不是错误 |
调 einsum 报错最容易犯的错是“对着报错信息猜”。更好的方法是有意识地缩小范围:先只传一个操作数测试,比如 np.einsum('ij->ji', A),确认下标本身没问题,再加入更多操作数。
6.2 我调试 einsum 的三步法
第一步,构造极小规模的已知数据。比如两个 2×2 矩阵,手算出目标结果,再用 einsum 跑一遍,对比是否一致。这个步骤能把“下标语义错误”卡在最早期,而不是等数据规模变大后才发现结果全错。
第二步,把 einsum 拆成等价的手写实现,输出中间步骤。例如怀疑 'bik,kj->bij' 写错,就手动做批量矩阵乘算一遍。如果结果一致,说明表达式是对的,问题在别处;如果不一致,把每个张量形状打印出来,确认下标对应的轴是不是你想的那个。对 PyTorch 用户,最直接的工具就是单步 debug + shape 打印。
第三步,对张量的每个维度起有意义的名称。我习惯在代码注释里写下 # x: (B, H, W, C) 这样的标注,再写 einsum 时下标一目了然。不要图省事全用 i, j, k 堆叠,后续维护时你会感谢自己留了下标注释。
6.3 一个必须注意的坑:省略号的兼容性差异
不同框架、不同版本对 einsum 省略号的支持程度是不同的。早期 NumPy 和部分科学计算库对省略号的处理比较保守,有些版本不支持批量维度下直接使用省略号。PyTorch 在旧版本里对省略号的解析也有不少边界 case 报错,后来才逐渐改进。
因此,如果你的代码要同时跑在多个环境里,尽量把 ... 显式写成具体字母,例如 'bik,bkj->bij' 代替 '...ik,...kj->...ij'。虽然省略号更通用,但显式字母让你完全掌握维度的排列,也能避免某些老版本库的兼容问题。
还有一个容易踩的:在显式模式下,如果输入张量是标量(零维),某些实现会直接不支持。碰到这种情况,先 reshape(1) 或 squeeze 一下再进 einsum。
6.4 如何让 einsum 代码更好维护
团队协作时,光有等价代码还不够,还要有可维护性。我总结几条经验:
- 每个 einsum 表达式必须有注释,至少写清各下标代表什么维度,比如
'bqhd,bkhd->bqhk' # b: batch, q: query, k: key, h: head, d: head_dim - 复杂的收缩表达式拆成多个小表达式,不要一口气写一个十几个字母的下标串。可读性优先,因为性能差异通常不大。
- 给 einsum 表达式封装成函数,并写单元测试,用随机张量对比参考实现。
- 不要用隐式模式,永远写满箭头右侧。
- 追求性能时,对比 baseline 后再决定是否用
optimize=True参数(部分库支持指定最优路径)。
其中第一条最重要。我记得有一次急着提交代码,写了个 'bij,bjk->bik' 没注释,过了两周自己看都费劲。后来强制自己所有 einsum 都写注释,效率反而提高了。
6.5 从调试到提升:自己写一个小型 eisum 解析器练手
如果想彻底吃透 einsum,这里提供一个进阶玩法:自己实现一个支持部分功能的迷你 einsum。基本思路是:
- 解析下标字符串,把字母映射到轴索引
- 根据输出下标,先对操作数做
transpose,把下标顺序调成计算需要的顺序 - 对需要收缩的维度,用
reshape+matmul替代求和 - 最后再
transpose到目标输出形状
自己动手写一遍,比看十篇博客都有效。过程中你会发现 einsum 的优化空间为什么存在:手动实现很难避免中间张量和重复的内存操作,而一个全局的调度器能在更高维度上做规划。我在学习阶段就写过一个小实现,虽然性能一般,但彻底搞懂了 'ik,kj->ij' 在底层到底发生了什么。
最后一点随想
实战里用了这么久 einsum,我最大的体会是,它不只是一种函数写法,更是一种思维方式。它逼着你把“一个张量的哪个维度对应什么含义”这个问题想清楚,而不是放任自己在 reshape 和 transpose 的缝隙里来回试探。刚开始接触 einsum 时你可能觉得下标难认,但一旦上手,你会发现自己写公式和写代码的界限在变模糊,很多复杂的线性代数表达直接在代码层面就能复现。慢慢练,保持好奇心,你的张量运算能力会上一个台阶。
