Сбор по индексам в PyTorch
Для каждой ячейки будущего результата
функция gather берёт значение
из исходного тензора: параметр
dim задаёт ось, вдоль которой
читают номера, а тензор индексов
имеет ту же форму, что и ответ.
Исходная таблица и набор номеров столбцов для каждого ряда:
import torch
source = torch.tensor([[10, 11, 12], [20, 21, 22]])
col_pick = torch.tensor([[0, 2], [1, 0]])
result = torch.gather(source, 1, col_pick)
print(source)
print(col_pick)
print(result) # выведет tensor([[10, 12], [21, 20]])
В первом ряду результата стоят
10 и 12 - нулевой
и второй столбцы исходника;
во втором - 21 и 20.
Сбор вдоль рядов устроен так же,
если номера задают строки:
import torch
source = torch.tensor([[10, 11], [20, 21], [30, 31]])
row_pick = torch.tensor([[2], [0]])
column = torch.gather(source, 0, row_pick)
print(column) # выведет tensor([[30], [10]])
Создайте таблицу [[100, 200, 300], [400, 500, 600]]
и таблицу номеров [[2, 0], [1, 2]]
той же формы; соберите значения
по столбцам и выведите ответ
2 на 2.
Создайте таблицу из трёх рядов
[1, 2], [3, 4], [5, 6]
и таблицу номеров рядов
[[0], [2]]; получите
столбец из двух чисел и выведите его.
Создайте ряд [7, 8, 9, 10]
и ряд номеров [3, 1, 0];
соберите три значения в том порядке,
как заданы номера, и выведите ряд.