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

Метод scatter

Метод scatter класса Tensor позволяет записывать значения из исходного тензора в целевой тензор по указанным индексам. Он принимает три основных параметра: dim - ось, вдоль которой производится индексация, index - тензор с индексами, и src - тензор с записываемыми значениями. Важно, чтобы размерности тензоров index и src совпадали.

Синтаксис

tensor.scatter(dim, index, src)

Также метод имеет версию с изменением на месте scatter_:

tensor.scatter_(dim, index, src)

Параметры метода:

  • dim - ось, вдоль которой будет происходить запись (целое число);
  • index - тензор с индексами, указывающими позиции для записи;
  • src - тензор, содержащий записываемые значения.

Пример

Выполним простую операцию записи: возьмем пустой тензор размером (1, 5) и запишем в него значения по индексам:

import torch t = torch.zeros(1, 5) index = torch.tensor([[0, 2, 4]]) src = torch.tensor([[10, 20, 30]]) res = t.scatter(1, index, src) print(res)

Результат выполнения кода:

tensor([[10., 0., 20., 0., 30.]])

Значения из src были записаны в тензор t по позициям, указанным в index. Обратите внимание, что запись идет по оси dim=1 (по столбцам).

Пример

Создадим one-hot векторы с помощью метода scatter:

import torch labels = torch.tensor([1, 0, 2]) num_classes = 3 one_hot = torch.zeros(labels.size(0), num_classes) one_hot.scatter_(1, labels.unsqueeze(1), 1) print(one_hot)

Результат выполнения кода:

tensor([ [0., 1., 0.], [1., 0., 0.], [0., 0., 1.] ])

Здесь мы преобразовали метки классов в one-hot представление. Метод scatter_ записал единицы в столбцы, соответствующие индексам из тензора labels.

Пример

Рассмотрим запись по другой оси. Заполним двумерный тензор по строкам (dim=0):

import torch t = torch.zeros(3, 2) index = torch.tensor([ [0, 1], [1, 0], [2, 2] ]) src = torch.tensor([ [5, 6], [7, 8], [9, 10] ]) t.scatter_(0, index, src) print(t)

Результат выполнения кода:

tensor([ [5., 8.], [7., 6.], [9., 10.] ])

В данном примере запись происходила вдоль оси строк (dim=0). Значения из src были распределены в соответствии с индексами строк из index.

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

  • метод gather,
    который выбирает значения из тензора по индексам
  • метод index_select,
    который возвращает элементы по указанным индексам
  • метод masked_fill,
    который заполняет элементы по маске
  • метод copy_,
    который копирует элементы из другого тензора
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить