PyTorch torch.gather Function
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)
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]])
Other Extensions