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