基于重复矩阵的爱因斯坦求和优化:机器学习矩阵计算性能提升
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:
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)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.
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

