torch.argwhere
-
torch.argwhere(input) → Tensor -
Возвращает тензор, содержащий индексы всех ненулевых элементов
input. Каждая строка в результате содержит индексы ненулевого элемента вinput. Результат отсортирован лексикографически, при этом последний индекс изменяется быстрее всего (стиль C).Если
inputимеет измерений, то результирующий тензор индексовoutимеет размер , где — общее количество ненулевых элементов в тензореinput.Примечание
Эта функция похожа на NumPy’s
argwhere.Когда
inputнаходится на CUDA, эта функция вызывает синхронизацию хост-устройство.- Параметры:
-
{input} –
Пример:
>>> t = torch.tensor([1, 0, 1]) >>> torch.argwhere(t) tensor([[0], [2]]) >>> t = torch.tensor([[1, 0, 1], [0, 1, 1]]) >>> torch.argwhere(t) tensor([[0, 0], [0, 2], [1, 1], [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.argwhere.html