PyTorch torch.take Function
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)
# 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)
# 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)
# 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)
# 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).
Other Extensions