Spec-Zone.ru › PyTorch 1

torch.argwhere

torch.argwhere(input) → Tensor

Возвращает тензор, содержащий индексы всех ненулевых элементов input. Каждая строка в результате содержит индексы ненулевого элемента в input. Результат отсортирован лексикографически, при этом последний индекс изменяется быстрее всего (стиль C).

Если input имеет nn измерений, то результирующий тензор индексов out имеет размер (z×n)(z \times n), где zz — общее количество ненулевых элементов в тензоре 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API