einsum是numpy对爱因斯坦求和约定的直接封装,性能取决于下标字符串合理性;漏写输出下标、误用广播、未启用optimize等均导致低效。

einsum 不是“高性能的爱因斯坦求和实现”,它本身就是 NumPy 对爱因斯坦求和约定的直接封装——性能好不好,取决于你写的下标字符串是否合理、是否避免了中间数组。
为什么 einsum 有时比 np.dot 或 np.matmul 慢?
很多人默认 einsum 天然更快,但实际它只是通用接口:底层会根据下标自动选择计算路径(如是否调用 BLAS),也可能退化为 Python 循环。比如 np.einsum('ij,jk->ik', a, b) 确实等价于 np.dot(a, b),但若写成 np.einsum('ij,jk', a, b)(漏掉箭头和输出下标),它会先算完整外积再收缩,内存暴涨且极慢。
- 必须显式写出输出下标,否则默认做隐式求和(即所有未在输出中出现的下标全求和)
- 含重复下标的输入项(如
'ii')会被自动对角化或求迹,不是矩阵乘法 - 当维度超过 3,
einsum的路径优化(optimize=True)才真正起作用;小规模运算反而可能因路径分析开销更慢
einsum 下标字符串怎么写才不翻车?
核心就一条:每个字母代表一个轴,相同字母表示该轴要配对相乘,未出现在输出中的字母表示沿该轴求和。常见错误是混淆广播与配对。
-
'i,i->':两个一维向量点积(结果标量),不是'i,i'(这会返回外积) -
'ijk,kl->ijl':类似 batched matrix multiply,其中k是求和轴,i,j,l是保留轴 -
'...ij,...jk->...ik':省略号支持前导批量维度,比手动reshape更安全,尤其处理(b, n, m)和(b, m, p)时 - 如果某轴在输入中出现奇数次(如
'ij,j'),它会在输出中保留——这是合法的,对应张量收缩而非矩阵乘
什么时候该用 optimize 参数?
optimize 控制是否启用路径优化,默认 'greedy',对 4+ 张量操作明显有用;但对两三个数组的小运算,设为 False 反而更快。
-
optimize=True会调用opt_einsum的逻辑,生成最优 contraction order,适合np.einsum('ab,bc,cd,de->ae', *four_matrices)这类链式乘法 -
optimize='optimal'极慢,仅调试用;生产环境用'greedy'或'dp'即可 - 首次调用时加
optimize=True会缓存路径,后续同一下标字符串复用,所以别怕第一次稍慢
替代 einsum 的更简单写法有哪些?
不是所有场景都值得写下标字符串。NumPy 已把高频模式封装成专用函数,它们通常更易读、更稳、且底层调用高度优化。
- 矩阵乘:
np.matmul(a, b)或a @ b,比einsum('ij,jk->ik', a, b)清晰且快 - 批量矩阵乘:
np.einsum('...ij,...jk->...ik', a, b)可直接用np.matmul(a, b)(只要 shape 兼容) - 对角线/迹:
np.diag(a)或np.trace(a),比einsum('ii->i', a)或einsum('ii', a)更明确 - 元素级乘积后求和:
np.sum(a * b, axis=1)比einsum('ij,ij->i', a, b)更直白,且现代 NumPy 对广播 +sum优化极好
最常被忽略的一点:下标字符串里字母顺序决定输出 shape 的轴序,而 einsum 不检查输入数组是否真的满足广播规则——它只按你写的下标机械配对。一旦维度对不上,报错信息是 ValueError: operands could not be broadcast together,但根本原因往往藏在下标和实际 shape 的错位里。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











