Метод masked_select
Метод masked_select класса Tensor возвращает одномерный тензор,
содержащий все элементы исходного тензора, для которых соответствующий
элемент маски имеет значение True. Первым параметром метод принимает
тензор-маску с битовыми значениями. Размерности тензора и маски должны
удовлетворять условиям вещания (broadcasting). Метод полезен для
фильтрации данных, удаления выбросов или выделения элементов,
удовлетворяющих определенному условию.
Синтаксис
tensor.masked_select(mask)
Пример
Давайте выберем из тензора элементы, которые больше числа 2:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = t > 2
res = t.masked_select(mask)
print(res)
Результат выполнения кода:
tensor([3, 4, 5])
Пример
Давайте используем маску с отрицанием для выбора элементов,
которые не равны 3:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = t != 3
res = t.masked_select(mask)
print(res)
Результат выполнения кода:
tensor([1, 2, 4, 5])
Пример
Давайте создадим маску на основе двумерного условия и применим ее к двумерному тензору:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
mask = (t % 2) == 0
res = t.masked_select(mask)
print(res)
Результат выполнения кода:
tensor([2, 4, 6])
Пример
Давайте используем маску меньшей размерности. Благодаря вещанию маска будет применена ко всем элементам тензора:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
mask = torch.tensor([True, False, True])
res = t.masked_select(mask)
print(res)
Результат выполнения кода:
tensor([1, 3, 4, 6])
Смотрите также
-
метод
index_select,
который выбирает элементы по индексам вдоль указанной размерности -
метод
masked_fill,
который заполняет элементы тензора по маске указанным значением -
метод
gather,
который собирает элементы по указанным индексам -
метод
scatter,
который записывает значения в тензор по указанным индексам