PyTorch torch.select Function


Pytorch torch 参考手册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)

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)

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)

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]。


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions