PyTorch torch.tril Function


Pytorch torch 参考手册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 tensor
  • diagonal: diagonal index, 0 indicates the main diagonal
  • out: 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)

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)

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)

The output is:

tensor([[1, 0, 0],
        [4, 5, 0]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions