You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何解读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:收缩后剩余的张量集合与最终目标下标

以给出的路径为例:

  1. 第一步ij,iiii->ij:将Y(下标ij)与X(下标iiii)按i维度对应相乘,得到形状为(N,N)的张量(每个元素为Y[i,j] * X[i,i,i,i])
  2. 第二步il,ik->ikl:将两个Y张量(下标il和ik)按i维度对应相乘,得到形状为(N,N,N)的张量(每个元素为Y[i,l] * Y[i,k])
  3. 第三步ij,im->ijm:将第一步的结果与Y(下标im)按i维度对应相乘,得到形状为(N,N,N)的张量(每个元素为第一步结果[i,j] * Y[i,m])
  4. 第四步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

优化点说明

  1. 简化对角线提取:用np.einsum('iiii->i', X)直接提取X[i,i,i,i],减少一次冗余的einsum调用
  2. 广播替代einsum:元素级对应相乘用numpy广播实现,避免einsum的下标解析开销,底层执行效率更高
  3. tensordot加速求和:np.tensordot自动调用BLAS库进行张量收缩,比手动einsum求和的速度更快

性能验证

该实现的性能会更接近einsum_two,且避免了optimize='optimal'的路径搜索开销:

  • 小N时:开销远小于einsum_two,性能接近原fast_einsum
  • 大N时:利用BLAS加速,性能显著优于原fast_einsum,与einsum_two持平

内容的提问来源于stack exchange,提问作者Solarflare0

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.24 12:37:56