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

基于重复矩阵的爱因斯坦求和优化:机器学习矩阵计算性能提升

Optimizing the Repeated Einsum Expression with Matrix Gram Matrices

Great question! Let's break down how we can simplify that repeated einsum by leveraging matrix operations—specifically Gram matrices and element-wise products—to drastically cut down computation time.

Step 1: Decompose the Original Einsum

First, let's unpack what your original einsum is doing:

np.einsum('ia,ib,ic,jb,jc,jd,je,kd,ke->ka', A, B, C, B, C, B, C, B, C)

Expanding the summation, this calculates:
$$\text{result}[k,a] = \sum_{i,b,c,j,d,e} A[i,a] \cdot B[i,b]C[i,c] \cdot B[j,b]C[j,c] \cdot B[j,d]C[j,e] \cdot B[k,d]C[k,e]$$

We can group terms by shared indices to find patterns:

  • The terms involving $b,c$ simplify to $\left(\sum_b B[i,b]B[j,b]\right) \cdot \left(\sum_c C[i,c]C[j,c]\right)$ — this is the product of the $(i,j)$ entries of $B$'s Gram matrix and $C$'s Gram matrix.
  • Similarly, the terms involving $d,e$ simplify to $\left(\sum_d B[j,d]B[k,d]\right) \cdot \left(\sum_e C[j,e]C[k,e]\right)$ — the product of the $(j,k)$ entries of $B$'s and $C$'s Gram matrices.

Step 2: Define Intermediate Matrices

Let's formalize these patterns into actionable steps:

  1. Gram Matrices: Compute the Gram matrices for B and C (these capture dot products of all row pairs in each matrix):

    B_gram = B @ B.T  # Equivalent to np.einsum('ib,jb->ij', B, B)
    C_gram = C @ C.T  # Equivalent to np.einsum('ic,jc->ij', C, C)
    
  2. Element-wise Product Matrix: Create matrix $M$ where each entry combines row similarity across both B and C:

    M = B_gram * C_gram  # Element-wise multiplication
    

    $M[i,j]$ represents how similar row $i$ and row $j$ are when considering both B and C together.

  3. Matrix Squaring: The sum over $j$ in the original expression is equivalent to multiplying $M$ by itself (matrix multiplication):

    M_sq = M @ M
    

    $M_{\text{sq}}[i,k]$ captures the cumulative similarity between row $i$ and row $k$ via all intermediate rows $j$.

Step 3: Compute the Final Result

Finally, multiply this squared matrix with $A$ to get the desired output:

result = M_sq.T @ A  # Equivalent to np.einsum('ik,ia->ka', M_sq, A)

Why This Is Faster

  • Original Complexity: The original einsum runs in $O(n^2 p^2 q^2)$ time, where $n$ is the number of rows in B/C, $p$ is the number of columns in B, and $q$ is the number of columns in C.
  • Optimized Complexity: The new approach runs in $O(n^2 p + n^2 q + n^3 + n^2 a)$ time (where $a$ is the number of columns in A). For large $p$ and $q$, this is an enormous improvement—often reducing computation from hours to seconds.

Verification

You can validate this with a small test case: generate random small matrices for A, B, C, compute both the original einsum and the optimized version, and check that the results are nearly identical (accounting for floating-point precision).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:58:00