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

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