Функция scatter_add
Функция scatter_add выполняет суммирование значений из тензора src
в тензор self по индексам, указанным в тензоре index.
В отличие от scatter, которая просто записывает значения,
scatter_add складывает новые значения с уже существующими.
Это полезно для агрегации данных, например, при обратном распространении
градиентов или построении гистограмм.
Синтаксис
torch.scatter_add(input, dim, index, src)
Параметры функции:
-
input- тензор, в который производится суммирование -
dim- ось (измерение), вдоль которой производится суммирование -
index- тензор с индексами, куда суммировать значения -
src- тензор с исходными значениями для суммирования
Все тензоры, кроме dim, должны быть одного размера
или иметь возможность широковещательного расширения.
Пример
Давайте создадим тензор и применим к нему функцию scatter_add:
import torch
t = torch.tensor([
[0, 0, 0],
[0, 0, 0],
])
indices = torch.tensor([
[0, 1, 0],
[1, 0, 1],
])
src = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.scatter_add(t, 1, indices, src)
print(res)
Результат выполнения кода:
tensor([
[4, 2, 0],
[5, 10, 0],
])
В результате значения из src суммируются в t
по индексам из indices вдоль первой оси.
Пример
Рассмотрим использование scatter_add в качестве метода тензора:
import torch
t = torch.tensor([
[10, 20, 30],
[40, 50, 60],
])
indices = torch.tensor([
[0, 2, 1],
[1, 0, 2],
])
src = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = t.scatter_add(1, indices, src)
print(res)
Результат выполнения кода:
tensor([
[11, 23, 22],
[45, 54, 66],
])
Функция добавляет значения из src к существующим элементам t
по указанным индексам.
Пример
Давайте используем scatter_add для построения гистограммы:
import torch
torch.manual_seed(0)
values = torch.randint(0, 5, (100,))
hist = torch.zeros(5, dtype=torch.int64)
indices = values.unsqueeze(0)
src = torch.ones_like(indices)
hist.scatter_add_(0, indices, src)
print(hist)
Результат выполнения кода:
tensor([22, 13, 18, 21, 26])
В результате мы получили гистограмму значений от 0 до 4.
Пример
Рассмотрим суммирование в трёхмерном тензоре:
import torch
t = torch.zeros(2, 3, 4)
indices = torch.tensor([
[[0, 1], [1, 0], [0, 1]],
[[1, 0], [0, 1], [1, 0]],
])
src = torch.ones(2, 3, 2)
res = torch.scatter_add(t, 2, indices, src)
print(res)
Результат выполнения кода:
tensor([
[
[2., 2., 0., 0.],
[2., 2., 0., 0.],
[2., 2., 0., 0.],
],
[
[2., 2., 0., 0.],
[2., 2., 0., 0.],
[2., 2., 0., 0.],
],
])
Значения суммируются в указанные индексы по последней оси.