Метод scatter_
Метод scatter_ класса Tensor выполняет запись значений в тензор по заданным индексам, изменяя сам тензор (in-place операция). Первым параметром передаётся размерность (dim), по которой производится запись. Вторым параметром передаётся индексный тензор (index), указывающий позиции для записи. Третьим параметром передаётся тензор значений (src) или скалярное значение. Индексный тензор должен иметь ту же размерность, что и исходный тензор, либо быть транслируемым к ней.
Синтаксис
tensor.scatter_(dim, index, src)
tensor.scatter_(dim, index, value)
Пример
Давайте создадим нулевой тензор и запишем в него значения из другого тензора по индексам:
import torch
t = torch.zeros(5)
idx = torch.tensor([0, 2, 4])
src = torch.tensor([10, 20, 30])
t.scatter_(0, idx, src)
print(t)
Результат выполнения кода:
tensor([10., 0., 20., 0., 30.])
Пример
Давайте заполним тензор одним и тем же значением по указанным индексам:
import torch
t = torch.zeros(4)
idx = torch.tensor([1, 3])
t.scatter_(0, idx, 99)
print(t)
Результат выполнения кода:
tensor([ 0., 99., 0., 99.])
Пример
Рассмотрим работу scatter_ с двумерным тензором по первой размерности:
import torch
t = torch.zeros(3, 4)
idx = torch.tensor([
[0, 1, 2, 0],
[2, 0, 1, 1],
])
src = torch.tensor([
[10, 20, 30, 40],
[50, 60, 70, 80],
])
t.scatter_(0, idx, src)
print(t)
Результат выполнения кода:
tensor([
[10., 0., 0., 40.],
[ 0., 20., 70., 80.],
[50., 0., 30., 0.],
])
Пример
Теперь выполним запись по второй размерности двумерного тензора:
import torch
t = torch.zeros(2, 4)
idx = torch.tensor([
[0, 2, 3, 1],
[1, 3, 0, 2],
])
src = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
t.scatter_(1, idx, src)
print(t)
Результат выполнения кода:
tensor([
[1., 4., 2., 3.],
[7., 5., 8., 6.],
])
Пример
Важно помнить, что если несколько значений записываются в один индекс, последнее записанное значение перезаписывает предыдущие:
import torch
t = torch.zeros(3)
idx = torch.tensor([1, 0, 1])
src = torch.tensor([10, 20, 30])
t.scatter_(0, idx, src)
print(t)
Результат выполнения кода:
tensor([20., 30., 0.])
Смотрите также
-
метод
scatter,
который выполняет ту же операцию, но возвращает новый тензор -
метод
gather,
который извлекает значения из тензора по индексам -
метод
index_select,
который выбирает значения по индексам из заданной размерности -
метод
masked_fill_,
который заполняет элементы тензора значением по маске