PyTorch torch.take Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.takeIt is a function in PyTorch used to retrieve elements at given index positions. It treats the input tensor as a one-dimensional array and returns the elements at the specified index positions.

Function Definition

torch.take(input, index)

Parameters:

  • input(Tensor): Input tensor.
  • index(Tensor): Integer index tensor specifying the positions of elements to retrieve.

Return Value:

  • torch.Tensor: Returns a new tensor consisting of elements at the specified index positions.

Usage Example

Example

import torch

# Create a 2D tensor
x = torch.tensor([[1, 2, 3],
                  [4, 5, 6],
                  [7, 8, 9]])

print("Original tensor:")
print(x)

# Treat the 2D tensor as a one-dimensional array, indices 0-8
# Get elements at indices 0, 4, 8
index = torch.tensor([0, 4, 8])
y = torch.take(x, index)

print("nIndex:", index)
print("Extracted elements:", y)

The output result is:

原始张量:
tensor([[1, 2, 3],
        [4, 5, 6],
        [7, 8, 9]])

索引: tensor([0, 4, 8])
取出的元素: tensor([1, 5, 9])

Example

import torch

# 3D tensor
x = torch.arange(24).reshape(2, 3, 4)

print("Original 3D tensor:")
print(x)
print("Shape:", x.shape)

# Get elements at multiple indices
index = torch.tensor([0, 1, 2, 10, 20, 23])
y = torch.take(x, index)

print("nIndex:", index)
print("Extracted elements:", y)

The output result is:

原始3D张量:
tensor([[[ 0,  1,  2,  3],
         [ 4,  5,  6,  7],
         [ 8, 9, 10, 11]],

        [[12, 13, 14, 15],
         [16, 17, 18, 19],
         [20, 21, 22, 23]]])
形状: torch.Size([2, 3, 4])

索引: tensor([ 0,  1,  2, 10, 20, 23])
取出的元素: tensor([ 0,  1,  2, 10, 20, 23])

Example

import torch

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

# Negative indices are counted from the end
index = torch.tensor([0, -1])  # The first and last elements
y = torch.take(x, index)

print("Original:", x)
print("Index [0, -1]:", y)

The output result is:

原始: tensor([[1, 2, 3],
        [4, 5, 6]])
索引 [0, -1]: tensor([1, 6])

Example

import torch

# Randomly select elements
x = torch.randn(10, 10)

# Randomly generate 10 indices
index = torch.randint(0, 100, (10,))
print("Random indices:", index)

# Retrieve elements at the corresponding positions
selected = torch.take(x, index)

print("Original tensor shape:", x.shape)
print("Shape of selected elements:", selected.shape)
print("Selected elements:", selected)

The output result is:

随机索引: tensor([12, 45, 67, 82, 35, 59, 92,  7, 28, 73])
原始张量形状: torch.Size([10, 10])
选中的元素形状: torch.Size([10])
选中的元素: tensor([ 0.2345, -0.1234,  0.5678,  1.2345, -0.6789,  0.8901, -0.3456,  0.1234, -0.5678,  0.7890])

Note:torch.takeThe input tensor is treated as a flattened one-dimensional tensor for indexing. Indices must be within the valid range (0 to numel-1). Negative indices can also be used (-1 represents the last element).


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions