Метод masked_fill
Метод masked_fill заменяет элементы тензора на заданное значение в тех позициях, где соответствующая маска имеет значение True. Первым параметром метод принимает маску (тензор булевого типа или целочисленный, который будет преобразован в булевый), вторым параметром - значение, которым нужно заполнить выбранные элементы. Метод не изменяет исходный тензор, а возвращает новый тензор с применёнными изменениями.
Синтаксис
t.masked_fill(mask, value)
Пример
Создадим тензор и заменим элементы, соответствующие маске, на значение -1:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = torch.tensor([False, True, False, True, False])
res = t.masked_fill(mask, -1)
print(res)
Результат выполнения кода:
tensor([ 1, -1, 3, -1, 5])
Пример
Применим маску к двумерному тензору, заменяя элементы на 0:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
mask = torch.tensor([
[True, False, True],
[False, True, False],
])
res = t.masked_fill(mask, 0)
print(res)
Результат выполнения кода:
tensor([
[0, 2, 0],
[4, 0, 6],
])
Пример
Маска может быть получена в результате сравнения элементов тензора с некоторым значением:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
mask = t > 25
res = t.masked_fill(mask, -1)
print(res)
Результат выполнения кода:
tensor([10, 20, -1, -1, -1])
Пример
Используем целочисленную маску (ненулевые значения интерпретируются как True):
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = torch.tensor([0, 1, 0, 1, 0])
res = t.masked_fill(mask, 99)
print(res)
Результат выполнения кода:
tensor([ 1, 99, 3, 99, 5])
Пример
Значение для заполнения может быть тензором, если оно скалярное или совместимо по форме с исходным тензором:
import torch
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
mask = torch.tensor([False, True, False, True, False])
res = t.masked_fill(mask, torch.tensor(0.0))
print(res)
Результат выполнения кода:
tensor([1., 0., 3., 0., 5.])
Смотрите также
-
метод
masked_fill_,
который выполняет заполнение на месте -
метод
masked_select,
который выбирает элементы тензора по маске -
метод
index_fill_,
который заполняет элементы по индексам -
метод
fill_,
который заполняет все элементы тензора значением