Метод masked_fill_
Метод masked_fill_ класса Tensor изменяет тензор,
заменяя значения элементов в тех позициях, где маска
содержит True, на переданное значение.
Этот метод работает непосредственно с исходным тензором,
то есть является операцией ⁅i⁆in-place⁅/i⁆.
Метод принимает два обязательных параметра: маску (тензор булевых значений или чисел) и значение, которым нужно заполнить отмеченные элементы.
Синтаксис
t.masked_fill_(mask, value)
Пример
Давайте создадим тензор и заполним его элементы,
превышающие определённое значение, числом -1:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = t > 3
t.masked_fill_(mask, -1)
print(t)
Результат выполнения кода:
tensor([1, 2, 3, -1, -1])
В данном примере маска mask содержит True
для элементов больше трёх. В этих позициях значения
были заменены на -1.
Пример
Создадим двумерный тензор и заменим все отрицательные элементы на ноль:
import torch
t = torch.tensor([
[1, -2, 3],
[-4, 5, -6],
])
mask = t < 0
t.masked_fill_(mask, 0)
print(t)
Результат выполнения кода:
tensor([
[1, 0, 3],
[0, 5, 0],
])
Все отрицательные числа в тензоре были заменены на нули.
Пример
Заполним элементы тензора по маске, созданной на основе другого тензора, с помощью сравнения:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
indices = torch.tensor([1, 3])
mask = torch.isin(t, indices)
t.masked_fill_(mask, 999)
print(t)
Результат выполнения кода:
tensor([10, 999, 30, 999, 50])
В данном примере мы создали маску, проверяющую,
принадлежат ли элементы тензора t указанным
индексам. В позициях, где условие истинно, значения
были заменены на 999.
Пример
Используем метод с маской, созданной случайным образом, и фиксируем результат для воспроизводимости:
import torch
torch.manual_seed(0)
t = torch.tensor([1.5, 2.7, 3.2, 4.1, 5.9])
mask = torch.rand(5) > 0.5
t.masked_fill_(mask, 0.0)
print(t)
Результат выполнения кода:
tensor([0.0000, 2.7000, 0.0000, 0.0000, 5.9000])
В этом примере половина элементов тензора была случайным образом заменена на ноль.
Смотрите также
-
метод
fill_,
который заполняет тензор одним значением -
метод
masked_select,
который возвращает элементы по маске -
метод
index_fill_,
который заполняет элементы по индексам -
метод
scatter_,
который заполняет элементы по указанным индексам