PyTorch torch.tensordot Function


Pytorch torch 参考手册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 tensor
  • other: The second input tensor
  • dims: 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)

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)

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)

The output result is:

torch.Size([2, 5])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions