PyTorch torch.gather Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.gatheris a function in PyTorch used to gather elements at the specified indices along the specified dimension.

Function Definition

torch.gather(input, dim, index, sparse_grad)

Usage Example

Example

import torch

x = torch.tensor([[1, 2], [3, 4], [5, 6]])

# Gather along dim=1
index = torch.tensor([[0], [1], [0]])
result = torch.gather(x, dim=1, index=index)

print(result)

The output is:

tensor([[1],
        [4],
        [5]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions