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

Метод 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_,
    который заполняет элементы тензора значением по маске
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить