PyTorch torch.diagonal Function


Pytorch torch 参考手册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 tensor
  • diagonal: Diagonal index, 0 represents the main diagonal
  • dim1: First dimension
  • dim2: 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)

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)

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)

The output result is:

torch.Size([2, 3])
tensor([[ 0,  5, 10],
        [12, 17, 22]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions