Python中numpy einsum的更快替代方案(dot/tensordot实现)
用
tensordot/广播替代einsum实现指定张量运算 场景1:einsum('ij, kjh->ikjh')的替代实现
原运算逻辑为 res1[i,k,j,h] = A[i,j] * B[k,j,h],本质是广播后逐元素相乘(无求和操作),tensordot核心是内积求和,因此不适用该场景。直接通过维度扩展+广播相乘即可高效实现:
import numpy as np A = np.arange(6).reshape((2, 3)) B = np.arange(12).reshape((2, 3, 2)) # 扩展A维度为(2,1,3,1),B维度为(1,2,3,2),广播后逐元素相乘 res1 = A[:, np.newaxis, :, np.newaxis] * B[np.newaxis, :, :, :]
验证:输出shape为(2,2,3,2),与原einsum结果完全一致。
场景2:einsum('ijk, jk->ij')的替代实现
原运算逻辑为 res2[i,j] = sum_k C[i,j,k] * D[j,k],是对k轴的内积求和,可用tensordot直接实现,关键是指定对齐求和的轴:
import numpy as np C = np.arange(12).reshape((2, 3, 2)) D = np.arange(6).reshape((3, 2)) # 对齐C的第2轴(k)与D的第1轴(k),执行内积求和 res2 = np.tensordot(C, D, axes=((2), (1)))
验证:输出结果为[[ 1 13 41] [ 7 43 95]],与原einsum结果完全一致。
补充说明
- 若需多轴对齐求和,
tensordot的axes参数可传入多轴索引列表,例如axes=((1,2), (0,1)); - 无求和的张量运算优先考虑广播+逐元素操作,比
einsum更高效。
内容的提问来源于stack exchange,提问作者Lyapunov1729
相关产品推荐
相关产品推荐

