PyTorch torch.diagonal Function
Pytorch torch Reference Manual
torch.diagonalIt is a function in PyTorch for extracting diagonal elements of a tensor. It returns a view of the specified diagonal of the input tensor.
Function Definition
torch.diagonal(input, diagonal=0, dim1=0, dim2=1)
Parameter Description:
input: Input tensordiagonal: Diagonal index, 0 represents the main diagonaldim1: First dimensiondim2: Second dimension
Usage Examples
Example
import torch
# Create a 3x3 matrix
x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the main diagonal
y = torch.diagonal(x)
print(y)
# Create a 3x3 matrix
x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the main diagonal
y = torch.diagonal(x)
print(y)
The output result is:
tensor([1, 5, 9])
Example
import torch
# Create a 3x3 matrix
x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the diagonal above the main diagonal
y = torch.diagonal(x, offset=1)
print(y)
# Create a 3x3 matrix
x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the diagonal above the main diagonal
y = torch.diagonal(x, offset=1)
print(y)
The output result is:
tensor([2, 6])
Example
import torch
# Create a 3D tensor
x = torch.arange(12).reshape(2, 3, 4)
# Extract the diagonal of the specified dimensions
y = torch.diagonal(x, offset=0, dim1=1, dim2=2)
print(y.shape)
print(y)
# Create a 3D tensor
x = torch.arange(12).reshape(2, 3, 4)
# Extract the diagonal of the specified dimensions
y = torch.diagonal(x, offset=0, dim1=1, dim2=2)
print(y.shape)
print(y)
The output result is:
torch.Size([2, 3])
tensor([[ 0, 5, 10],
[12, 17, 22]])
Other Extensions