PyTorch torch.diag Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.diagIs a function in PyTorch used to create diagonal matrices or extract diagonal elements of a tensor.

Function Definition

torch.diag(input, diagonal=0, out=None)

Parameter Description:

  • input: Input tensor
  • diagonal: Diagonal index, 0 represents the main diagonal, positive values represent upper diagonals, negative values represent lower diagonals

Usage Example

Example

import torch

# Create a one-dimensional tensor
x = torch.tensor([1, 2, 3])

# Create a diagonal matrix
y = torch.diag(x)
print(y)

The output result is:

tensor([[1, 0, 0],
        [0, 2, 0],
        [0, 0, 3]])

Example

import torch

# Extract diagonal elements from a matrix
x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])

# Extract the main diagonal
y = torch.diag(x)
print(y)

The output result is:

tensor([1, 5, 9])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions