PyTorch torch.einsum Function
PyTorch torch Reference Manual
torch.einsumis a function in PyTorch used for the Einstein summation convention, which can concisely express various tensor operations.
Function Definition
torch.einsum(equation, *operands)
Usage Examples
Examples
import torch
# Matrix transpose
A = torch.randn(2, 3)
AT = torch.einsum('ij->ji', A)
print("Transpose shape:", AT.shape)
# Matrix multiplication
B = torch.randn(3, 4)
C = torch.einsum('ij,jk->ik', A, B)
print("Multiplication shape:", C.shape)
# Dot product
a = torch.randn(3)
b = torch.randn(3)
dot = torch.einsum('i,i->', a, b)
print("Dot product:", dot.item())
# Matrix transpose
A = torch.randn(2, 3)
AT = torch.einsum('ij->ji', A)
print("Transpose shape:", AT.shape)
# Matrix multiplication
B = torch.randn(3, 4)
C = torch.einsum('ij,jk->ik', A, B)
print("Multiplication shape:", C.shape)
# Dot product
a = torch.randn(3)
b = torch.randn(3)
dot = torch.einsum('i,i->', a, b)
print("Dot product:", dot.item())
Other Extensions