PyTorch torch.select_scatter Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.select_scatteris a function in PyTorch used to scatter the source tensor to a specified index position. Itsrcscatters the values ofinputinto the specified dimensiondimand indexindexposition.

Function Definition

torch.select_scatter(input, src, dim, index)

Parameters:

  • input(Tensor): The input tensor, i.e., the tensor to be modified.
  • src(Tensor): The source tensor, the values to be scattered into input.
  • dim(int): The dimension of scattering.
  • index(int): The position index of scattering.

Return Value:

  • torch.Tensor: Returns the modified tensor.

Usage Examples

Example

import torch

# Create the input tensor and source tensor
input = torch.zeros(4, 4)
src = torch.ones(4)

# Scatter src at the index 1 position in the first dimension (rows)
output = torch.select_scatter(input, src, dim=0, index=1)

print("Input tensor:")
print(input)
print("nSource tensor:")
print(src)
print("nScatter result at index 1 position:")
print(output)

The output result is:

输入张量:
tensor([[0., 0., 0., 0.],
        [0., 0., 0., 0.],
        [0., 0., 0., 0.],
        [0., 0., 0., 0.]])

源张量:
tensor([1., 1., 1., 1.])

在索引1位置散布结果:
tensor([[0., 0., 0., 0.],
        [1., 1., 1., 1.],
        [0., 0., 0., 0.],
        [0., 0., 0., 0.]])

Example

import torch

# Create a 3D tensor
input = torch.zeros(3, 4, 5)
src = torch.ones(5)

# Scatter in the second dimension (index 2)
output = torch.select_scatter(input, src, dim=1, index=2)

print("Input shape:", input.shape)
print("Source shape:", src.shape)
print("Result shape:", output.shape)
print("nResult (dim=1, index=2):")
print(output[0])

The output result is:

输入形状: torch.Size([3, 4, 5])
源形状: torch.Size([5])
结果形状: torch.Size([3, 4, 5])

结果 (dim=1, index=2):
tensor([1., 1., 1., 1., 1.])

Example

import torch

# Use different values
input = torch.arange(16).reshape(4, 4).float()
src = torch.tensor([100, 100, 100, 100])

# Scatter at the last row
output = torch.select_scatter(input, src, dim=0, index=-1)

print("Original tensor:")
print(input)
print("nScattered to index -1 position:")
print(output)

The output result is:

原始张量:
tensor([[ 0.,  1.,  2.,  3.],
        [ 4.,  5.,  6.,  7.],
        [ 8.,  9., 10., 11.],
        [12., 13., 14., 15.]])

散布到索引-1位置:
tensor([[ 0.,  1.,  2.,  3.],
        [ 4.,  5.,  6.,  7.],
        [ 8.,  9., 10., 11.],
        [100., 100., 100., 100.]])

Note:torch.select_scatterDoes not modify the original input tensor, but returns a new tensor.srcThe dimensions must match the dimensions of the slice to be replaced.


Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions