PyTorch torch.where function
Pytorch torch reference manual
torch.whereIt is a function in PyTorch used to return elements based on conditions.
Function definition
torch.where(condition, input, other)
Usage example
Example
import torch
condition = torch.tensor([[True, False], [False, True]])
x = torch.tensor([[1, 2], [3, 4]])
y = torch.tensor([[10, 20], [30, 40]])
result = torch.where(condition, x, y)
print(result)
condition = torch.tensor([[True, False], [False, True]])
x = torch.tensor([[1, 2], [3, 4]])
y = torch.tensor([[10, 20], [30, 40]])
result = torch.where(condition, x, y)
print(result)
The output result is:
tensor([[ 1, 20],
[30, 4]])
Other extensions