Функция scatter
Функция scatter выполняет запись значений в тензор по заданным индексам.
Первый параметр dim определяет ось, вдоль которой производится запись.
Второй параметр index содержит индексы для записи.
Третий параметр src содержит записываемые значения.
Четвёртый параметр reduce определяет операцию агрегации (суммирование или перемножение).
Синтаксис
torch.scatter(dim, index, src)
torch.scatter(dim, index, src, reduce='add')
Функция возвращает новый тензор с записанными значениями. Исходный тензор не изменяется.
Пример
Давайте выполним простую запись значений по индексам:
import torch
t = torch.tensor([[1, 2, 3], [4, 5, 6]])
index = torch.tensor([[0, 1], [1, 0]])
src = torch.tensor([[9, 8], [7, 6]])
res = torch.scatter(t, 0, index, src)
print(res)
Результат выполнения кода:
tensor([
[9, 8],
[7, 6],
])
Пример
Давайте выполним запись с агрегацией по сумме:
import torch
t = torch.tensor([[1, 2], [3, 4]])
index = torch.tensor([[0, 1], [1, 0]])
src = torch.tensor([[5, 6], [7, 8]])
res = torch.scatter(t, 1, index, src, reduce='add')
print(res)
Результат выполнения кода:
tensor([
[6, 8],
[11, 7],
])
Пример
Давайте выполним запись с агрегацией по произведению:
import torch
t = torch.tensor([[2, 3], [4, 5]])
index = torch.tensor([[0, 1], [1, 0]])
src = torch.tensor([[1, 2], [3, 4]])
res = torch.scatter(t, 0, index, src, reduce='multiply')
print(res)
Результат выполнения кода:
tensor([
[2, 12],
[12, 5],
])
Смотрите также
-
функцию
scatter_add,
которая выполняет атомарное суммирование по индексам -
функцию
gather,
которая собирает значения из тензора по индексам -
функцию
index_select,
которая выбирает элементы по индексам -
функцию
take,
которая извлекает элементы по индексам из плоского тензора