如何使用np.einsum实现指定轴参数的np.tensordot张量点积操作
Replacing np.tensordot with np.einsum for Your 4D Arrays
To replicate the np.tensordot(a, b, axes=([1,2,3], [0,1,2])) operation using np.einsum, we just need to map the contracted axes to shared indices in the einsum notation. Here's how to do it:
Step-by-Step Explanation
- Assign indices to each axis:
- For array
a(shape(A0, A1, A2, A3)), label its axes asi, j, k, l(whereicorresponds to axis 0, andj/k/lcorrespond to axes 1/2/3). - For array
b(shape(B0, B1, B2, B3)), since we're contracting its axes 0/1/2 witha's 1/2/3, label those matching axes asj, k, l, and axis 3 asm.
- For array
- Write the einsum string:
- The string
'ijkl,jklm->im'tells einsum to sum over the shared indicesj, k, l(the contracted axes), leaving onlyi(froma's axis 0) andm(fromb's axis 3) as the output axes.
- The string
Code Example
Let’s verify with sample arrays to ensure the results match:
import numpy as np # Create random 4D arrays with compatible shapes a = np.random.rand(2, 3, 4, 5) # Shape (2,3,4,5) b = np.random.rand(3, 4, 5, 6) # Shape (3,4,5,6) — contracted axes match a's 1-3 # Original tensordot result tensordot_result = np.tensordot(a, b, axes=([1,2,3], [0,1,2])) # Equivalent einsum result einsum_result = np.einsum('ijkl,jklm->im', a, b) # Check if results are identical (within floating-point tolerance) print(np.allclose(tensordot_result, einsum_result)) # Output: True
Why This Works
Both operations perform the exact same contraction: they sum the product of elements where the contracted axes (a's 1-3 and b's 0-2) align. The einsum notation just makes this axis mapping explicit with indices, which can be more readable once you’re familiar with the syntax.
内容的提问来源于stack exchange,提问作者kazar4
相关产品推荐
相关产品推荐

