PyTorch torch.trace Function
Pytorch torch Reference Manual
torch.traceis a function in PyTorch used to compute the trace of a matrix. The trace is the sum of the elements on the main diagonal of a matrix.
Function Definition
torch.trace(input)
Parameter Description:
input: input tensor (at least two-dimensional)
Usage Example
Example
import torch
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Compute the trace (sum of main diagonal elements)
y = torch.trace(a)
print(y)
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Compute the trace (sum of main diagonal elements)
y = torch.trace(a)
print(y)
The output result is:
tensor(15)
Example
import torch
# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Compute the trace
y = torch.trace(a)
print(y)
# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Compute the trace
y = torch.trace(a)
print(y)
The output result is:
tensor(6)
Example
import torch
# Create a multi-dimensional tensor, only consider the first two dimensions
a = torch.randn(3, 4, 4, 5)
y = torch.trace(a)
print(y.shape)
# Create a multi-dimensional tensor, only consider the first two dimensions
a = torch.randn(3, 4, 4, 5)
y = torch.trace(a)
print(y.shape)
The output result is:
torch.Size([3, 5])
Other Extensions