PyTorch torch.tril_indices Function


Pytorch torch 参考手册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 rows
  • column: Number of columns
  • offset: Diagonal offset
  • dtype: Data type of the returned value
  • device: 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)

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)

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)

The output result is:

row: tensor([0, 1, 1, 2, 2, 2])
col: tensor([0, 0, 1, 0, 1, 2])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions