РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
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 для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить