Метод 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_,
который копирует элементы из другого тензора