PyTorch torch.tril_indices Function
PyTorch torch Reference Manual
torch.tril_indicesIt is a function in PyTorch used to generate lower triangular matrix indices. It returns the row indices and column indices of the lower triangular portion (including the diagonal).
Function Definition
torch.tril_indices(row, column, offset=0, dtype=torch.long, device='cpu')
Parameter Description:
row: Number of rowscolumn: Number of columnsoffset: Diagonal offsetdtype: Data type of the returned valuedevice: Device
Usage Example
Example
import torch
# Generate lower triangular indices for a 3x3 matrix
row, col = torch.tril_indices(3, 3)
print("row:", row)
print("col:", col)
# Generate lower triangular indices for a 3x3 matrix
row, col = torch.tril_indices(3, 3)
print("row:", row)
print("col:", col)
The output result is:
row: tensor([0, 1, 1, 2, 2, 2]) col: tensor([0, 0, 1, 0, 1, 2])
Example
import torch
# Generate indices and use them for indexing operations
row, col = torch.tril_indices(3, 3, offset=1)
# Create a 3x3 matrix
a = torch.ones(3, 3)
# Use the indices to set the lower triangular part
a[row, col] = 0
print(a)
# Generate indices and use them for indexing operations
row, col = torch.tril_indices(3, 3, offset=1)
# Create a 3x3 matrix
a = torch.ones(3, 3)
# Use the indices to set the lower triangular part
a[row, col] = 0
print(a)
The output result is:
tensor([[1., 0., 0.],
[1., 1., 0.],
[1., 1., 1.]])
Example
import torch
# Non-square matrix case
row, col = torch.tril_indices(3, 4)
print("row:", row)
print("col:", col)
# Non-square matrix case
row, col = torch.tril_indices(3, 4)
print("row:", row)
print("col:", col)
The output result is:
row: tensor([0, 1, 1, 2, 2, 2]) col: tensor([0, 0, 1, 0, 1, 2])
Other Extensions