PyTorch torch.index_select Function


Pytorch torch 参考手册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)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions