Функция masked_fill
Метод masked_fill заменяет элементы тензора на заданное значение
в тех позициях, где соответствующая маска имеет значение True.
Маска должна быть той же размерности, что и тензор, либо
поддерживать вещание (broadcasting). Метод возвращает новый тензор
и не изменяет исходный.
Синтаксис
tensor.masked_fill(mask, value)
Пример
Давайте заменим все отрицательные числа в тензоре на ноль:
import torch
t = torch.tensor([1, -2, 3, -4, 5])
mask = t < 0
res = t.masked_fill(mask, 0)
print(res)
Результат выполнения кода:
tensor([1, 0, 3, 0, 5])
Пример
Теперь заполним элементы, которые больше или равны трём, значением 99:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = t >= 3
res = t.masked_fill(mask, 99)
print(res)
Результат выполнения кода:
tensor([1, 2, 99, 99, 99])
Пример
Рассмотрим работу с двумерным тензором и маской:
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, -1)
print(res)
Результат выполнения кода:
tensor([
[-1, 2, -1],
[4, -1, 6],
])
Пример
Маска может быть результатом сравнения двух тензоров:
import torch
t = torch.tensor([10, 20, 30, 40])
threshold = torch.tensor([15, 25, 35, 45])
mask = t < threshold
res = t.masked_fill(mask, 0)
print(res)
Результат выполнения кода:
tensor([0, 0, 0, 40])
Смотрите также
-
функцию
masked_select,
которая выбирает элементы тензора по маске -
функцию
index_fill,
которая заполняет элементы по индексам -
функцию
where,
которая выбирает элементы из двух тензоров по условию -
функцию
zeros_like,
которая создаёт тензор из нулей по образцу