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