PyTorch torch.index_select Function
Pytorch torch Reference Manual
torch.index_selectIt is a function in PyTorch used to select elements corresponding to indices along a specified dimension.
Function Definition
torch.index_select(input, dim, index)
Usage Example
Example
import torch
x = torch.randn(4, 5)
# Select rows 0 and 2
indices = torch.tensor([0, 2])
result = torch.index_select(x, dim=0, index=indices)
print("Result shape:", result.shape)
x = torch.randn(4, 5)
# Select rows 0 and 2
indices = torch.tensor([0, 2])
result = torch.index_select(x, dim=0, index=indices)
print("Result shape:", result.shape)
Other Extensions