Функция index_fill
Функция index_fill заполняет элементы тензора
по заданным индексам определённым значением.
Первый параметр принимает тензор, который нужно изменить.
Второй параметр - размерность, по которой производится
заполнение. Третий параметр - тензор с индексами.
Четвёртый параметр - значение, которым нужно заполнить.
Синтаксис
torch.index_fill(input, dim, index, value)
Метод вызывается от тензора:
t.index_fill_(dim, index, value)
Пример с одномерным тензором
Давайте заполним элементы с индексами 1 и 3
значением 99:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
indices = torch.tensor([1, 3])
res = torch.index_fill(t, 0, indices, 99)
print(res)
Результат выполнения кода:
tensor([ 1, 99, 3, 99, 5])
Пример с in-place версией
Используем метод index_fill_ для изменения
исходного тензора без создания нового:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
indices = torch.tensor([0, 2, 4])
t.index_fill_(0, indices, 0)
print(t)
Результат выполнения кода:
tensor([ 0, 20, 0, 40, 0])
Пример с двумерным тензором
Заполним элементы по столбцам 0 и 2
во всех строках значением -1:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
indices = torch.tensor([0, 2])
res = torch.index_fill(t, 1, indices, -1)
print(res)
Результат выполнения кода:
tensor([
[-1, 2, -1],
[-1, 5, -1],
[-1, 8, -1],
])
Пример с многомерным тензором
Заполним элементы по последней размерности для трёхмерного тензора:
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
])
indices = torch.tensor([0, 1])
res = torch.index_fill(t, 2, indices, 100)
print(res)
Результат выполнения кода:
tensor([
[
[100, 100],
[100, 100],
],
[
[100, 100],
[100, 100],
],
])
Смотрите также
-
функцию
index_select,
которая возвращает элементы тензора по индексам -
функцию
masked_fill,
которая заполняет элементы по маске -
функцию
scatter,
которая записывает значения по заданным индексам -
функцию
gather,
которая собирает значения по индексам