PyTorch torch.diag Function
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 tensordiagonal: 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)
# 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)
# 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])
Other Extensions