Метод gather
Метод gather класса Tensor собирает значения из исходного тензора
по индексам, указанным в другом тензоре. Первым параметром метод принимает
целое число - ось (dim), вдоль которой производится сборка.
Вторым параметром передается тензор индексов (index).
Третьим параметром можно указать флаг sparse_grad,
который управляет режимом разреженного градиента.
Синтаксис
t.gather(dim, index, sparse_grad=False)
Пример
Давайте создадим двумерный тензор и соберем значения по столбцам:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
indices = torch.tensor([
[0, 1, 0],
[1, 0, 1],
])
res = t.gather(dim=0, index=indices)
print(res)
Результат выполнения кода:
tensor([
[1, 5, 3],
[4, 2, 6],
])
Пример
Теперь соберем значения по строкам:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
indices = torch.tensor([
[2, 0, 1],
[1, 2, 0],
])
res = t.gather(dim=1, index=indices)
print(res)
Результат выполнения кода:
tensor([
[3, 1, 2],
[5, 6, 4],
])
Пример
Давайте используем метод для выборки значений из трехмерного тензора:
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
])
indices = torch.tensor([
[
[0, 1],
[1, 0],
],
])
res = t.gather(dim=0, index=indices)
print(res)
Результат выполнения кода:
tensor([
[
[1, 6],
[7, 4],
],
])
Пример
При использовании метода важно, чтобы тензор индексов имел ту же размерность, что и исходный тензор, за исключением оси, по которой производится сборка:
import torch
t = torch.tensor([
[10, 20, 30],
[40, 50, 60],
])
indices = torch.tensor([
[1, 0],
])
res = t.gather(dim=1, index=indices)
print(res)
Результат выполнения кода:
tensor([
[20, 10],
[50, 40],
])
Смотрите также
-
метод
index_select,
который выбирает значения по индексам вдоль оси -
метод
masked_select,
который выбирает значения по маске -
метод
scatter,
который записывает значения по индексам -
метод
scatter_,
который изменяет тензор по индексам