PyTorch torch.tensordot Function
PyTorch torch Reference Manual
torch.tensordotIt is a function in PyTorch used to compute the dot product of two tensors along specified dimensions. It is a Tensor Contraction operation.
Function Definition
torch.tensordot(input, other, dims=2)
Parameter Description:
input: The first input tensorother: The second input tensordims: The number of dimensions to contract or a list of dimension pairs
Usage Example
Example
import torch
# Create two one-dimensional tensors
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Compute the dot product
y = torch.tensordot(a, b, dims=1)
print(y)
# Create two one-dimensional tensors
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
# Compute the dot product
y = torch.tensordot(a, b, dims=1)
print(y)
The output result is:
tensor(32)
Example
import torch
# Create two two-dimensional tensors
a = torch.tensor([[1, 2], [3, 4]])
b = torch.tensor([[5, 6], [7, 8]])
# Compute matrix multiplication (contract two dimensions)
y = torch.tensordot(a, b, dims=2)
print(y)
# Create two two-dimensional tensors
a = torch.tensor([[1, 2], [3, 4]])
b = torch.tensor([[5, 6], [7, 8]])
# Compute matrix multiplication (contract two dimensions)
y = torch.tensordot(a, b, dims=2)
print(y)
The output result is:
tensor(70)
Example
import torch
# Create a three-dimensional tensor
a = torch.randn(2, 3, 4)
b = torch.randn(3, 4, 5)
# Specify the dimension pairs to contract
y = torch.tensordot(a, b, dims=[[1, 2], [0, 1]])
print(y.shape)
# Create a three-dimensional tensor
a = torch.randn(2, 3, 4)
b = torch.randn(3, 4, 5)
# Specify the dimension pairs to contract
y = torch.tensordot(a, b, dims=[[1, 2], [0, 1]])
print(y.shape)
The output result is:
torch.Size([2, 5])
Other Extensions