РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
136 of 769 menu

Функция 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.], ], ])

Значения суммируются в указанные индексы по последней оси.

Смотрите также

  • функцию scatter,
    которая записывает значения по индексам вместо суммирования
  • функцию gather,
    которая собирает значения по индексам
  • функцию index_add,
    которая выполняет суммирование по одномерным индексам
  • функцию masked_scatter,
    которая записывает значения по маске
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить