torch.tril_indices
-
torch.tril_indices(row, col, offset=0, *, dtype=torch.long, device='cpu', layout=torch.strided) → Tensor -
Возвращает индексы нижней треугольной части матрицы
rowнаcol, в тензоре 2xN, где первая строка содержит координаты строк всех индексов, а вторая строка содержит координаты столбцов. Индексы упорядочены по строкам, а затем по столбцам.Нижняя треугольная часть матрицы определяется как элементы на и ниже диагонали.
Аргумент
offsetуправляет тем, какую диагональ учитывать. Еслиoffset= 0, сохраняются все элементы на и ниже главной диагонали. Положительное значение включает столько же диагоналей выше главной диагонали, а отрицательное значение исключает столько же диагоналей ниже главной диагонали. Главная диагональ — это множество индексов для , где — размеры матрицы.Примечание
При выполнении на CUDA,
row * colдолжно быть меньше , чтобы предотвратить переполнение при вычислении.- Параметры:
-
-
row (
int) – количество строк в матрице 2-D. -
col (
int) – количество столбцов в матрице 2-D. -
offset (
int) – смещение диагонали относительно главной диагонали. По умолчанию: если не указано, 0.
-
row (
- Ключевые аргументы:
-
-
dtype (
torch.dtype, необязательно) – желаемый тип данных возвращаемого тензора. По умолчанию: еслиNone,torch.long. -
device (
torch.device, необязательно) – желаемое устройство возвращаемого тензора. По умолчанию: еслиNone, используется текущее устройство для типа тензора по умолчанию (см.torch.set_default_tensor_type()).deviceбудет CPU для типов тензоров CPU и текущим устройством CUDA для типов тензоров CUDA. -
layout (
torch.layout, необязательно) – в настоящее время поддерживается толькоtorch.strided.
-
dtype (
Пример:
>>> a = torch.tril_indices(3, 3) >>> a tensor([[0, 1, 1, 2, 2, 2], [0, 0, 1, 0, 1, 2]]) >>> a = torch.tril_indices(4, 3, -1) >>> a tensor([[1, 2, 2, 3, 3, 3], [0, 0, 1, 0, 1, 2]]) >>> a = torch.tril_indices(4, 3, 1) >>> a tensor([[0, 0, 1, 1, 1, 2, 2, 2, 3, 3, 3], [0, 1, 0, 1, 2, 0, 1, 2, 0, 1, 2]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.tril_indices.html