Функция argwhere
Функция argwhere возвращает тензор, содержащий индексы всех ненулевых элементов входного тензора.
Результат представляет собой двумерный тензор, где каждая строка содержит координаты одного ненулевого элемента.
Первый параметр функции - входной тензор. Вторым параметром можно передать флаг as_tuple.
Синтаксис
torch.argwhere(input, [as_tuple])
Пример
Давайте найдем индексы ненулевых элементов в одномерном тензоре:
import torch
t = torch.tensor([0, 1, 0, 2, 3, 0])
res = torch.argwhere(t)
print(res)
Результат выполнения кода:
tensor([
[1],
[3],
[4],
])
Пример
Найдем индексы ненулевых элементов в двумерном тензоре:
import torch
t = torch.tensor([
[0, 1, 0],
[2, 0, 3],
[0, 0, 4],
])
res = torch.argwhere(t)
print(res)
Результат выполнения кода:
tensor([
[0, 1],
[1, 0],
[1, 2],
[2, 2],
])
Пример
Используем параметр as_tuple для получения кортежа тензоров индексов по каждому измерению:
import torch
t = torch.tensor([
[0, 1, 0],
[2, 0, 3],
[0, 0, 4],
])
res = torch.argwhere(t, as_tuple=True)
print(res)
Результат выполнения кода:
(tensor([0, 1, 1, 2]), tensor([1, 0, 2, 2]))
Пример
Применим argwhere к тензору с отрицательными значениями и нулями:
import torch
t = torch.tensor([-1, 0, 5, 0, -3, 2])
res = torch.argwhere(t)
print(res)
Результат выполнения кода:
tensor([
[0],
[2],
[4],
[5],
])
Смотрите также
-
функцию
nonzero,
которая возвращает индексы ненулевых элементов -
функцию
where,
которая возвращает элементы по условию -
функцию
masked_select,
которая выбирает элементы по маске -
функцию
index_select,
которая выбирает элементы по индексам