PyTorch torch.select Function
Pytorch torch Reference Manual
torch.selectIt is a function in PyTorch used to select the slice corresponding to an index along a specified dimension. It returns a sliced view at the specified index position on the specified dimension.
Function Definition
torch.select(input, dim, index)
Parameters:
input(Tensor): Input tensor.dim(int): The dimension to select.index(int): The index to select.
Return Value:
torch.Tensor: Returns the slice at the specified index position (dimension reduced by 1).
Usage Examples
Example
import torch
# Create a 3x4 tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Select the row with index 1 on the first dimension (rows)
y = torch.select(x, dim=0, index=1)
print("Original tensor:")
print(x)
print("nSelect the row with index 1:")
print(y)
# Create a 3x4 tensor
x = torch.tensor([[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]])
# Select the row with index 1 on the first dimension (rows)
y = torch.select(x, dim=0, index=1)
print("Original tensor:")
print(x)
print("nSelect the row with index 1:")
print(y)
The output result is:
原始张量:
tensor([[ 1, 2, 3, 4],
[ 5, 6, 7, 8],
[ 9, 10, 11, 12]])
选择索引1的行:
tensor([5, 6, 7, 8])
Example
import torch
# Create a 3D tensor
x = torch.arange(24).reshape(2, 3, 4)
print("Original 3D tensor:")
print(x)
print("Shape:", x.shape)
# Select the element with index 0 in the first dimension (batch)
y = torch.select(x, dim=0, index=0)
print("nSelect dim=0, index=0:")
print(y)
print("Shape:", y.shape)
# Create a 3D tensor
x = torch.arange(24).reshape(2, 3, 4)
print("Original 3D tensor:")
print(x)
print("Shape:", x.shape)
# Select the element with index 0 in the first dimension (batch)
y = torch.select(x, dim=0, index=0)
print("nSelect dim=0, index=0:")
print(y)
print("Shape:", y.shape)
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])
选择 dim=0, index=0:
tensor([[ 0, 1, 2, 3],
[ 4, 5, 6, 7],
[ 8, 9, 10, 11]])
形状: torch.Size([3, 4])
Example
import torch
# Use negative index
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# Equivalent way to select the last row
y = torch.select(x, dim=0, index=-1)
print("Original tensor:")
print(x)
print("nSelect the last row (index=-1):")
print(y)
# Use negative index
x = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# Equivalent way to select the last row
y = torch.select(x, dim=0, index=-1)
print("Original tensor:")
print(x)
print("nSelect the last row (index=-1):")
print(y)
The output result is:
原始张量:
tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
选择最后一行 (index=-1):
tensor([7, 8, 9])
Note:torch.selectIt returns a view, not a copy, so the operation is efficient. Similar functionality can also be achieved using slicing operations, such asx[index]orx[index:index+1]。
Other Extensions