einsum不报错需满足三条硬规则:输入下标长度等于张量维数;重复字母表示求和;输出字母必须在输入中至少出现一次。

einsum 怎么写才能不报错:下标字符串的硬规则
绝大多数 np.einsum() 报错,都卡在下标字符串格式不对——不是语法错,是逻辑冲突。比如你想把形状为 (2, 3, 4) 的张量按前两维缩并,写成 "ijk->k" 是对的;但写成 "ij->k" 就会直接抛 ValueError: einstein sum subscripts string contains too many subscripts for operand。
关键约束有三条:
- 每个输入 operand 的下标长度必须等于其维度数(
a.shape == (2,3,4)→ 下标必须是三位,如"ijk") - 重复出现的下标字母表示该轴要被求和(
"iij->j"中i出现两次 → 对第 0、1 维求和) - 输出下标中没出现的字母,自动被求和;输出中出现的字母,必须在所有输入中至少出现一次(
"ij,jk->ik"合法;"ij,jk->il"非法,因为l没定义)
替代 np.tensordot() 和 np.matmul() 的典型写法
很多人用 einsum 是为了绕过 tensordot 里 axes= 参数的绕口令式写法,或者避免 matmul 对高维张量的隐式广播规则踩坑。
常见映射关系:
-
np.tensordot(a, b, axes=([1], [0]))→np.einsum("ik,kj->ij", a, b) -
np.matmul(a, b)(当a.shape=(..., m, k),b.shape=(..., k, n))→np.einsum("...mk,...kn->...mn", a, b) - 三维 batch 矩阵乘:
np.einsum("nij,njk->nik", A, B)比np.matmul(A, B)更直白,且不依赖...广播的“省略号对齐”行为
性能陷阱:为什么有时候 einsum 比 np.sum() + np.multiply() 还慢
einsum 不是万能加速器。它底层调用的是优化过的 BLAS 或自生成循环,但优化效果高度依赖下标模式和数据规模。
容易慢的几种情况:
- 下标含太多自由变量(如
"ijkl,mnop->ijmnop"),导致中间结果爆炸,内存带宽成为瓶颈 - 使用
optimize=True时,小数组(optimize=False 更稳 - 等价于简单操作却写复杂下标:比如
np.einsum("i,i->", a, b)(点积)比np.dot(a, b)慢 2–3 倍;此时用原生函数更合适
高级索引场景:用 einsum 实现“条件缩并”或“动态轴选择”
标准 sum 或 tensordot 要求轴号写死,而 einsum 可以配合字符串拼接实现运行时决定缩并方式,适合写通用工具函数。
例如:给定一个张量 x 和要保留的轴名列表 keep = ["i", "k"],自动生成下标:
axes = "ijklm"
keep_axes = "ik"
input_subs = axes[:x.ndim]
output_subs = keep_axes
sum_axes = "".join(c for c in input_subs if c not in output_subs)
subscript = f"{input_subs}->{output_subs}"
result = np.einsum(subscript, x)
注意:这种动态拼接只适合调试或配置驱动场景;线上高频调用务必预编译好下标字符串,避免重复解析开销。
最常被忽略的一点:einsum 对输入数组的内存布局敏感。如果传入非 C-contiguous 数组(比如切片后未 .copy()),某些下标模式会退化成纯 Python 循环,速度骤降。遇到性能异常,先检查 a.flags.c_contiguous。










