爱因斯坦求和约定(Einstein Summation Convention)及einsum() 函数

爱因斯坦求和约定(Einstein Summation Convention)及einsum() 函数

一、爱因斯坦求和约定

爱因斯坦在写相对论张量公式时,发明的一套下标简写规则:

重复出现的下标,默认代表对这个维度求和 ,省去手写求和符号∑\sum∑。

是张量计算的下标简写规则。

核心规则:在同一项表达式里,下标如果重复出现一次(一对),代表对这个下标遍历全部取值并求和 ;只出现一次的下标,保留在结果中。
重复下标叫
哑标(dummy index):只用于求和,结果里消失;
只出现一次的下标叫
自由标(free index)**:留在输出,决定输出维度。
硬性规则:一个下标最多只能成对出现,不能出现3次及以上 ,否则非法。

例:AijBjkA_{ij}B_{jk}AijBjk 合法;AiiBikA_{ii}B_{ik}AiiBik 合法;AijBjjCjA_{ij}B_{jj}C_{j}AijBjjCj 下标j出现3次,不允许。
原始数学约定默认只对重复下标求和,不写∑\sum∑符号 。

PyTorch/numpy einsum 是这套规则的工程实现,增加了 -> 显式声明输出下标(数学公式里没有->,是计算机扩展语法)。函数名字 einsum = Ein + sum → Einstein + summation。


二、基础例子,从低维到高维

例1:向量点积(1维张量)

向量 a∈Rn,  b∈Rna\in\mathbb R^n,\;b\in\mathbb R^na∈Rn,b∈Rn

公式:s=aibis = a_i b_is=aibi

按照约定,i重复,对i求和:

s=∑i=1naibis=\sum_{i=1}^n a_i b_is=i=1∑naibi

python 复制代码
import torch
a = torch.tensor([1,2,3])
b = torch.tensor([4,5,6])
# i是哑标,求和,输出是标量
res = torch.einsum('i,i->', a, b)
print(res) # 1*4 +2*5 +3*6 = 32

例2:矩阵-向量乘法

Aij  (n×m),  xj(m)A_{ij}\; (n\times m),\; x_j(m)Aij(n×m),xj(m)

yi=Aijxjy_i = A_{ij} x_jyi=Aijxj

j成对出现,沿j求和;i只出现一次,保留为输出下标。

python 复制代码
A = torch.tensor([[1,2],[3,4]])
x = torch.tensor([10,20])
y = torch.einsum('ij,j->i', A, x)
print(y) # [1*10+2*20, 3*10+4*20] = [50,110]

例3:矩阵乘法(最经典)

Aik⋅Bkj=CijA_{ik} \cdot B_{kj}=C_{ij}Aik⋅Bkj=Cij

k是哑标,求和;i,j是自由标,保留。

Cij=∑kAikBkjC_{ij}=\sum_k A_{ik}B_{kj}Cij=k∑AikBkj

python 复制代码
A = torch.tensor([[1,2],[3,4]]) # i,k
B = torch.tensor([[5,6],[7,8]]) # k,j
C = torch.einsum('ik,kj->ij', A,B)
print(C)
# [[1*5+2*7, 1*6+2*8],
#  [3*5+4*7, 3*6+4*8]]
# [[19,22],[43,50]]

例4:矩阵内缩,矩阵迹 trace(对角线求和)

AiiA_{ii}Aii:下标i重复,对i求和

tr(A)=∑iAiitr(A)=\sum_i A_{ii}tr(A)=i∑Aii

python 复制代码
A = torch.tensor([[1,2],[3,4]])
res = torch.einsum('ii->', A)
print(res) #1+4=5

例5:外积(向量外积,不做求和)

Cij=aibjC_{ij}=a_i b_jCij=aibj,i,j都只出现一次,没有哑标,不求和

python 复制代码
a = torch.tensor([1,2])
b = torch.tensor([3,4])
C = torch.einsum('i,j->ij',a,b)
print(C)
# [[1*3,1*4],
#  [2*3,2*4]]

三、扩展到批量多维张量(NLP深度学习场景)

例6: SimpleAttention 里面的一行:

python 复制代码
att = torch.einsum('abc,c->ab', (inputs, self.word_weight))
  • inputs:abc = [batch, seq_len, hidden_dim]
  • word_weight:c = [hidden_dim]
    下标c成对,沿hidden_dim求和。
    输出下标ab → [batch, seq_len]
    含义:每个样本、每个词,词向量与weight向量做点积。
    等价:
python 复制代码
att = torch.sum(inputs * word_weight, dim=-1)

例7:MMA MatchTensor 核心双线性项

python 复制代码
matching_matrix = torch.einsum('abd,fde,ace->afbc', [x1, self.M, x2])
  • x1: abd → (batch, seqA, dimA)
  • M: fde → (channel, dimA, dimB)
  • x2: ace → (batch, seqB, dimB)
    哑标:d,e,沿着d、e两个维度求和收缩
    自由标:a,f,b,c,输出afbc = (batch, channel, seqA, seqB)

含义:对每个样本、每个通道,计算句子A每个词 和句子B每个词 的双线性匹配分数,得到seqA × seqB匹配矩阵。

如果用bmm+transpose写,代码会非常长,可读性差。

四、其它

箭头 -> 的作用(einsum独有,原始爱因斯坦公式没有)

原始数学写法,自动推断输出下标:只保留只出现一次的下标,按字母顺序输出 。

但深度学习里,我们需要显式控制输出维度顺序 ,所以pytorch/numpy的einsum增加->:

输入下标串 -> 输出下标串

举例子对比:

'ij,jk' 不写箭头,自动输出ik

'ij,jk->ki',手动指定输出下标顺序为k,i,相当于做完矩阵乘法再转置。

适用场景(深度学习)

✅ 适合:

  • 2维:矩阵乘法、向量点积、迹、外积
  • 3维:批量矩阵乘法、LSTM输出加权求和(SimpleAttention)
  • 4维张量交互:句子对匹配的词-词交互矩阵(MMA模型,你上面代码)

❌ 不适合:

超大规模张量,部分场景einsum底层优化不如专门算子(bmm, matmul)速度快;简单运算优先用专用算子。

注:

torch.matmul:通用矩阵乘法,自动适配维度,支持广播 ;

torch.bmm:批量矩阵乘法(Batch Matrix Multiply) ,严格3维张量,不支持广播。

五、PyTorch / NumPy / TensorFlow 的 einsum 对比

一句话结论:三者都提供 einsum,核心下标语法完全一样;但底层实现、参数、支持特性、边界行为有差别。

数学的爱因斯坦求和下标规则,是统一标准,所以同样的下标字符串 'ik,kj->ij' 在三个库里面含义相同。

1. 基础接口对照

NumPy
python 复制代码
np.einsum(subscripts, *operands, out=None, dtype=None, order='K', casting='safe')
  • 最老牌实现。
  • 输入是numpy数组。
  • 支持 -> 显式输出下标(新版numpy才完善支持)。
  • 可以指定输出数组out、控制数据类型。
PyTorch
python 复制代码
torch.einsum(equation, *operands)
  • 只有两个核心参数:方程字符串 + 张量列表。
  • 不支持 out/dtype 这类numpy那些高级参数。
  • 支持自动微分(可以放进神经网络,作为可计算图一部分,就是你MMA、Attention代码里用的)。
  • 支持CUDA张量,GPU上跑。
TensorFlow
python 复制代码
tf.einsum(equation, *inputs, name=None)
  • 接口风格接近PyTorch。
  • 支持tf的计算图、GPU、自动微分。
  • TF2里einsum是原生支持。

2. 相同点

  1. 下标语法完全一致 :'abc,c->ab' 三个库写出来等价,哑标/自由标规则一样。
  2. 都支持多维张量收缩、矩阵乘、外积、迹等操作。
  3. 都支持显式箭头 -> 指定输出维度顺序。

3. 关键区别

① 自动微分能力
  • NumPy np.einsum:不能求导!
    numpy数组不是计算图,纯数值运算,无梯度,绝对不能直接放到神经网络前向传播里。
  • PyTorch / TensorFlow einsum:支持自动微分
    张量记录计算历史,反向传播可以算梯度,所以你前面MMA、SimpleAttention才可以用torch.einsum。

👉 这就是深度学习代码只用torch/tf的einsum,不用numpy的根本原因。

② 性能与优化
  • NumPy einsum:CPU实现,老版本优化一般;高维大张量有时比较慢。
  • PyTorch einsum :内部会尝试把表达式优化、转成matmul/bmm等专用内核 。
    但复杂多下标4维收缩(像MMA的abd,fde,ace->afbc),不一定能完全优化,极端情况比手写bmm慢。GPU可用。
  • TensorFlow einsum:会交给XLA编译器优化,复杂表达式有时性能更好。

小提醒:非常简单的矩阵乘法,优先直接用@ / matmul,而不是einsum,专用算子通常更快。

③ 下标字母限制 & 特殊规则
  • Numpy:早期版本不强制要求->,可以省略输出下标,自动推断。
  • PyTorch:支持省略->自动推断输出,但工程代码强烈建议写->,可读性高。
  • TF:同样支持省略箭头。

统一最佳实践:写深度学习代码,一律带上 ->,方便读代码。

④ 广播行为差异

  • Numpy einsum:遵循numpy广播规则。
  • PyTorch einsum:遵循PyTorch广播。
    大部分场景感受不到差别;但维度不匹配时报错信息不一样。
⑤ 多输入张量数量

三个输入张量

python 复制代码
torch.einsum('abd,fde,ace->afbc', [x1, self.M, x2])

✅ PyTorch ✅ Numpy ✅ TF 全都支持多个输入张量收缩,不是只能两个张量相乘。

4. 代码对比例子,同一个式子三库写法

python 复制代码
# 向量点积
import numpy as np
a_np = np.array([1,2,3])
b_np = np.array([4,5,6])
res_np = np.einsum('i,i->', a_np, b_np)

import torch
a_torch = torch.tensor([1,2,3])
b_torch = torch.tensor([4,5,6])
res_torch = torch.einsum('i,i->', a_torch, b_torch)

import tensorflow as tf
a_tf = tf.constant([1,2,3])
b_tf = tf.constant([4,5,6])
res_tf = tf.einsum('i,i->', a_tf, b_tf)

计算结果数值完全一样。

5. 工程选型总结

  1. 离线数据分析(纯CPU,不用训练网络) :用np.einsum
  2. PyTorch深度学习网络(可微分、GPU) :torch.einsum
  3. TensorFlow/Keras模型 :tf.einsum
相关推荐
计算机毕业编程指导师几秒前
【计算机毕设推荐】基于Hadoop+Django的LLM多维度性能评估分析系统源码 毕业设计 选题推荐 毕设选题 数据分析 机器学习 深度学习
大数据·hadoop·python·django·毕业设计·课程设计·llm性能
AC赳赳老秦4 分钟前
OpenClaw 与 FineBI 联动方案:公开数据自动采集与实时业务分析看板实践
大数据·开发语言·python·php·finebi·deepseek·openclaw
计算机毕业编程指导师15 分钟前
【Python毕设选题推荐】基于Hadoop+Django的奥斯卡奖获奖数据可视化分析系统源码 毕业设计 选题推荐 毕设选题 数据分析 机器学习 深度学习
大数据·hadoop·python·计算机·毕业设计·课程设计·奥斯卡奖
风早爽太16 分钟前
Python 学习笔记:数据库迁移工具 ‌Alembic
数据库·python·fastapi·alembic
外收内放19 分钟前
06 | 优化篇① 游戏改名《鳞光纪》了:这次翻新,把“脸“全换了
python·游戏·pygame
计算机毕业编程指导师22 分钟前
【大数据毕设选题】基于Spark的电信网络诈骗话术语义特征挖掘分析系统源码 毕业设计 选题推荐 数据分析 机器学习 深度学习
大数据·hadoop·python·spark·毕业设计·课程设计·网络诈骗
for_ever_love__27 分钟前
电商订单数据清洗实战——Pandas 处理缺失值、重复单与异常金额
python·数据分析·pandas·数据清洗
计算机毕业编程指导师28 分钟前
计算机大数据毕设怎么选?基于Spark的个体肥胖健康风险评估的数据分析与可视化系统带你通关 源码 毕业设计 选题推荐 毕设选题 数据分析 机器学习
大数据·hadoop·python·计算机·spark·毕业设计·肥胖风险
坤坤子吖28 分钟前
Python基础语法学习:函数
开发语言·笔记·python·学习
玖石书28 分钟前
linux 安装uv(python虚拟环境管理)
linux·python·uv