Функция masked_select
Функция masked_select возвращает новый одномерный тензор,
содержащий элементы исходного тензора, для которых соответствующие
элементы маски имеют значение True. Первым параметром
функция принимает исходный тензор, вторым параметром - тензор маски
того же размера или транслируемый до размера исходного тензора.
Маска должна содержать логические значения (bool).
Синтаксис
torch.masked_select(input, mask)
Пример
Давайте выберем из тензора только те элементы, которые больше числа 3:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
mask = t > 3
res = torch.masked_select(t, mask)
print(res)
Результат выполнения кода:
tensor([4, 5])
Пример
Давайте выберем элементы из двумерного тензора, которые меньше или равны числу 4:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
mask = t <= 4
res = torch.masked_select(t, mask)
print(res)
Результат выполнения кода:
tensor([1, 2, 3, 4])
Обратите внимание, что результат всегда одномерный тензор, независимо от исходной размерности.
Пример
Давайте используем маску, созданную вручную, для выбора элементов из тензора:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
mask = torch.tensor([True, False, True, False, True])
res = torch.masked_select(t, mask)
print(res)
Результат выполнения кода:
tensor([10, 30, 50])
Пример
Давайте выберем элементы из тензора с плавающей точкой по сложному условию:
import torch
t = torch.tensor([0.5, 1.5, 2.5, 3.5, 4.5])
mask = (t > 1.0) & (t < 4.0)
res = torch.masked_select(t, mask)
print(res)
Результат выполнения кода:
tensor([1.5000, 2.5000, 3.5000])
Смотрите также
-
функцию
where,
которая выбирает элементы из двух тензоров по условию -
функцию
nonzero,
которая возвращает индексы ненулевых элементов -
функцию
argwhere,
которая возвращает координаты элементов, удовлетворяющих условию -
функцию
index_select,
которая выбирает элементы тензора по индексам