Функция 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,
которая выбирает элементы по маске логических значений