如何解读np.einsum_path输出,实现无优化开销的高效np.einsum调用
优化np.einsum调用:解读einsum_path并手动实现最优路径
问题背景
现有两个numpy数组:维度为(N,N,N,N)的X和维度为(N,N)的Y,目标是高效执行以下einsum操作:
Z = np.einsum('iiii,ij,ik,il,im->jklm', X, Y, Y, Y, Y)
测试了三种实现方式,输出结果一致但性能不同:
def einsum_one(X, Y): return np.einsum('iiii,ij,ik,il,im->jklm', X, Y, Y, Y, Y) def einsum_two(X, Y): return np.einsum('iiii,ij,ik,il,im->jklm', X, Y, Y, Y, Y, optimize='optimal') def fast_einsum(X, Y): Z = np.einsum('ij,iiii->ij', Y, X) W = np.einsum('il,ik->ikl', Y, Y) Z = np.einsum('ij,im->ijm', Z, Y) Z = np.einsum('ijm,ikl->jklm', Z, W) return Z
其中fast_einsum是根据np.einsum_path的最优路径输出手动实现的,路径输出如下:
Complete contraction: iiii,ij,ik,il,im->jklm Naive scaling: 5 Optimized scaling: 5 Naive FLOP count: 1.638e+05 Optimized FLOP count: 6.662e+04 Theoretical speedup: 2.459 Largest intermediate: 4.096e+03 elements -------------------------------------------------------------------------- scaling current remaining -------------------------------------------------------------------------- 2 ij,iiii->ij ik,il,im,ij->jklm 3 il,ik->ikl im,ij,ikl->jklm 3 ij,im->ijm ikl,ijm->jklm 5 ijm,ikl->jklm jklm->jklm
预期fast_einsum在N较大时最快,einsum_one最慢;einsum_two的优化开销在N小时明显,N大时可忽略。但基准测试结果(单位:微秒)显示fast_einsum并非最优:
N | einsum_one | einsum_two | fast_einsum 4 15.1 949 13.8 8 256 966 87.9 12 1920 1100 683 16 7460 1180 2490 20 21800 1390 7080
疑问:如何正确解读np.einsum_path的输出,在不使用optimize='optimal'的前提下,实现最快的einsum调用?
正确解读einsum_path输出
np.einsum_path的输出展示了最优的收缩步骤,每一行的含义如下:
- scaling:当前收缩操作的时间复杂度量级(如scaling 2对应O(N²),scaling 3对应O(N³))
- current:当前要执行的收缩操作,格式为
输入下标1,输入下标2->输出下标,表示将两个输入张量按指定下标收缩,得到新张量 - remaining:收缩后剩余的张量集合与最终目标下标
以给出的路径为例:
- 第一步
ij,iiii->ij:将Y(下标ij)与X(下标iiii)按i维度对应相乘,得到形状为(N,N)的张量(每个元素为Y[i,j] * X[i,i,i,i]) - 第二步
il,ik->ikl:将两个Y张量(下标il和ik)按i维度对应相乘,得到形状为(N,N,N)的张量(每个元素为Y[i,l] * Y[i,k]) - 第三步
ij,im->ijm:将第一步的结果与Y(下标im)按i维度对应相乘,得到形状为(N,N,N)的张量(每个元素为第一步结果[i,j] * Y[i,m]) - 第四步
ijm,ikl->jklm:将第三步和第二步的结果按i维度求和,得到最终的四阶张量Z
优化后的手动实现
原fast_einsum多次调用np.einsum带来了额外开销,且未利用numpy的广播、BLAS加速等底层优化。以下是更高效的实现:
def optimized_einsum(X, Y): # 提取X的对角线元素X[i,i,i,i] X_diag = np.einsum('iiii->i', X) # 第一步:Y的每一行乘以X_diag[i],用广播替代einsum Y_weighted = Y * X_diag[:, np.newaxis] # 第二步:Y[:,k]与Y[:,l]按i对应相乘,广播实现 Y_kl = Y[:, :, np.newaxis] * Y[:, np.newaxis, :] # shape (N,N,N) -> (i,k,l) # 第三步:Y_weighted[:,j]与Y[:,m]按i对应相乘,广播实现 Y_jm = Y_weighted[:, :, np.newaxis] * Y[:, np.newaxis, :] # shape (N,N,N) -> (i,j,m) # 第四步:对i维度求和,用tensordot利用BLAS加速 # tensordot后得到形状(j,m,k,l),转置为目标的(j,k,l,m) Z = np.tensordot(Y_jm, Y_kl, axes=([0], [0])).transpose(0, 2, 3, 1) return Z
优化点说明
- 简化对角线提取:用
np.einsum('iiii->i', X)直接提取X[i,i,i,i],减少一次冗余的einsum调用 - 广播替代einsum:元素级对应相乘用numpy广播实现,避免einsum的下标解析开销,底层执行效率更高
- tensordot加速求和:
np.tensordot自动调用BLAS库进行张量收缩,比手动einsum求和的速度更快
性能验证
该实现的性能会更接近einsum_two,且避免了optimize='optimal'的路径搜索开销:
- 小N时:开销远小于
einsum_two,性能接近原fast_einsum - 大N时:利用BLAS加速,性能显著优于原
fast_einsum,与einsum_two持平
内容的提问来源于stack exchange,提问作者Solarflare0
相关产品推荐
相关产品推荐

