PyTorch torch.masked_select Function
PyTorch torch Reference Manual
torch.masked_selectIt is a function in PyTorch used to select elements based on a boolean mask.
Function Definition
torch.masked_select(input, mask)
Usage Example
Example
import torch
x = torch.tensor([1, 2, 3, 4, 5])
mask = torch.tensor([True, False, True, False, True])
# Select elements based on the mask
result = torch.masked_select(x, mask)
print(result)
x = torch.tensor([1, 2, 3, 4, 5])
mask = torch.tensor([True, False, True, False, True])
# Select elements based on the mask
result = torch.masked_select(x, mask)
print(result)
The output result is:
tensor([1, 3, 5])
Other Extensions