PyTorch torch.triu Function
Pytorch torch Reference Manual
torch.triuIt is a function in PyTorch used to extract the upper triangular part of a matrix (including the main diagonal). The lower triangular part will be set to 0.
Function Definition
torch.triu(input, diagonal=0, out=None)
Parameter Description:
input: input tensordiagonal: diagonal index, 0 indicates the main diagonalout: output tensor
Usage Examples
Example
import torch
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the upper triangular part
y = torch.triu(a)
print(y)
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the upper triangular part
y = torch.triu(a)
print(y)
The output result is:
tensor([[1, 2, 3],
[0, 5, 6],
[0, 0, 9]])
Example
import torch
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the upper triangular part above the first diagonal above the main diagonal
y = torch.triu(a, diagonal=1)
print(y)
# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
# Extract the upper triangular part above the first diagonal above the main diagonal
y = torch.triu(a, diagonal=1)
print(y)
The output result is:
tensor([[0, 2, 3],
[0, 0, 6],
[0, 0, 0]])
Example
import torch
# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Extract the upper triangular part
y = torch.triu(a)
print(y)
# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Extract the upper triangular part
y = torch.triu(a)
print(y)
The output result is:
tensor([[1, 2, 3],
[0, 5, 6]])
Other Extensions