Функция where
Функция where возвращает тензор, состоящий из элементов,
выбранных из двух тензоров на основе условия. Первым параметром
функция принимает условие в виде булева тензора, вторым -
тензор для значений, где условие истинно, а третьим - тензор
для значений, где условие ложно. Если передан только один
параметр, функция возвращает индексы ненулевых элементов.
Синтаксис
torch.where(condition, x, y)
Пример
Давайте создадим два тензора и выберем элементы на основе условия:
import torch
condition = torch.tensor([True, False, True, False])
x = torch.tensor([1, 2, 3, 4])
y = torch.tensor([5, 6, 7, 8])
res = torch.where(condition, x, y)
print(res)
Результат выполнения кода:
tensor([1, 6, 3, 8])
Пример
Давайте используем функцию с условием сравнения элементов тензора:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = torch.where(t > 3, t, torch.tensor(0))
print(res)
Результат выполнения кода:
tensor([0, 0, 0, 4, 5])
Пример
Давайте применим функцию с двумерными тензорами:
import torch
condition = torch.tensor([
[True, False],
[False, True]
])
x = torch.tensor([
[1, 2],
[3, 4]
])
y = torch.tensor([
[5, 6],
[7, 8]
])
res = torch.where(condition, x, y)
print(res)
Результат выполнения кода:
tensor([
[1, 6],
[7, 4]
])
Пример
Давайте используем функцию с одним параметром для поиска индексов:
import torch
t = torch.tensor([0, 1, 0, 2, 0, 3])
res = torch.where(t != 0)
print(res)
Результат выполнения кода:
(tensor([1, 3, 5]),)
Смотрите также
-
функцию
nonzero,
которая возвращает индексы ненулевых элементов тензора -
функцию
argwhere,
которая возвращает индексы элементов, удовлетворяющих условию -
функцию
masked_select,
которая выбирает элементы по маске -
функцию
masked_fill,
которая заполняет элементы по маске заданным значением