Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
128 of 769 menu
◀ ▶

Функция 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,
    которая выбирает элементы по индексам
← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить