torch.where
-
torch.where(condition, x, y) → Tensor -
Возвращает тензор элементов, выбранных из
xилиy, в зависимости отcondition.Операция определяется следующим образом:
Примечание
Тензоры
condition,x,yдолжны быть совместимы по правилам векторизации.- Параметры:
-
- condition (BoolTensor) – При True (ненулевое значение), возвращается x, иначе возвращается y
-
x (Tensor или Скаляр) – значение (если
xявляется скаляром) или значения, выбранные в индексах, гдеconditionпринимает значениеTrue -
y (Tensor или Скаляр) – значение (если
yявляется скаляром) или значения, выбранные в индексах, гдеconditionпринимает значениеFalse
- Возвращает:
-
Тензор с формой, равной результату векторизации
condition,x,y - Тип возвращаемого значения:
Пример:
>>> x = torch.randn(3, 2) >>> y = torch.ones(3, 2) >>> x tensor([[-0.4620, 0.3139], [ 0.3898, -0.7197], [ 0.0478, -0.1657]]) >>> torch.where(x > 0, x, y) tensor([[ 1.0000, 0.3139], [ 0.3898, 1.0000], [ 0.0478, 1.0000]]) >>> x = torch.randn(2, 2, dtype=torch.double) >>> x tensor([[ 1.0779, 0.0383], [-0.8785, -1.1089]], dtype=torch.float64) >>> torch.where(x > 0, x, 0.) tensor([[1.0779, 0.0383], [0.0000, 0.0000]], dtype=torch.float64)- torch.where(condition) кортеж LongTensor
torch.where(condition)идентичноtorch.nonzero(condition, as_tuple=True).Примечание
См. также
torch.nonzero().
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.where.html