РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
130 of 769 menu

Функция gather

Функция gather выполняет сборку значений из исходного тензора по указанным индексам. Первым параметром передаётся исходный тензор input, вторым - измерение dim, по которому осуществляется выборка, третьим - тензор индексов index. Результат имеет ту же форму, что и тензор индексов.

Синтаксис

torch.gather(input, dim, index)

Пример

Давайте выберем элементы из одномерного тензора по указанным индексам:

import torch t = torch.tensor([10, 20, 30, 40, 50]) indices = torch.tensor([0, 2, 4]) res = torch.gather(t, 0, indices) print(res)

Результат выполнения кода:

tensor([10, 30, 50])

Пример

Теперь выполним сборку по строкам для двумерного тензора:

import torch t = torch.tensor([ [1, 2, 3], [4, 5, 6], [7, 8, 9], ]) indices = torch.tensor([ [0, 1, 2], [2, 0, 1], ]) res = torch.gather(t, 0, indices) print(res)

Результат выполнения кода:

tensor([ [1, 2, 3], [7, 5, 6], ])

Пример

Выполним сборку по столбцам для двумерного тензора:

import torch t = torch.tensor([ [1, 2, 3], [4, 5, 6], [7, 8, 9], ]) indices = torch.tensor([ [0, 2], [1, 0], ]) res = torch.gather(t, 1, indices) print(res)

Результат выполнения кода:

tensor([ [1, 3], [5, 4], ])

Смотрите также

  • функцию scatter,
    которая выполняет обратную операцию - запись значений по индексам
  • функцию index_select,
    которая выбирает элементы тензора по индексам вдоль измерения
  • функцию take,
    которая выбирает элементы из тензора по плоским индексам
  • функцию masked_select,
    которая выбирает элементы по маске логических значений
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить